Compare commits

..
Author SHA1 Message Date
geohot ccd753e1aa set testpath on pytest 2025-09-15 14:29:30 +08:00
376 changed files with 3401 additions and 42334 deletions
+57 -70
View File
@@ -28,7 +28,7 @@ jobs:
# since sudo is required for usbgpu on macos, move the cache to a new location, as some of the files are owned by root
PYTHONPYCACHEPREFIX: /tmp/tiny_python_pycache
runs-on: [self-hosted, macOS]
timeout-minutes: 60
timeout-minutes: 20
defaults:
run:
shell: bash -e -o pipefail {0}
@@ -52,16 +52,14 @@ jobs:
- name: reset process replay
run: python3.11 test/external/process_replay/reset.py
- name: Run Stable Diffusion
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
run: BENCHMARK_LOG=stable_diffusion JIT=1 ASSERT_MIN_STEP_TIME=500 python3.11 examples/stable_diffusion.py --fp16 --seed 0 --noshow --timing | tee sd.txt
- name: Run Stable Diffusion without fp16
run: BENCHMARK_LOG=stable_diffusion_fp32 JIT=1 ASSERT_MIN_STEP_TIME=900 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=700 python3.11 examples/stable_diffusion.py --seed 0 --noshow --timing | tee sd_no_fp16.txt
- name: Run Stable Diffusion v2
# TODO: very slow step time
run: BENCHMARK_LOG=stable_diffusion_v2 JIT=1 ASSERT_MIN_STEP_TIME=10000 python3.11 examples/sdv2.py --fp16 --seed 0 --noshow --timing | tee sdv2.txt
run: BENCHMARK_LOG=stable_diffusion_v2 JIT=1 ASSERT_MIN_STEP_TIME=1600 python3.11 examples/sdv2.py --fp16 --seed 0 --noshow --timing | tee sdv2.txt
# process replay can't capture this, the graph is too large
# TODO: too slow
# - name: Run SDXL
# run: BENCHMARK_LOG=stable_diffusion_xl ASSERT_MIN_STEP_TIME=5000 CAPTURE_PROCESS_REPLAY=0 JIT=1 python3.11 examples/sdxl.py --seed 0 --noshow --timing | tee sdxl.txt
- name: Run SDXL
run: BENCHMARK_LOG=stable_diffusion_xl ASSERT_MIN_STEP_TIME=3000 CAPTURE_PROCESS_REPLAY=0 JIT=1 python3.11 examples/sdxl.py --seed 0 --noshow --timing | tee sdxl.txt
- name: Run model inference benchmark
run: METAL=1 python3.11 test/external/external_model_benchmark.py
- name: Test speed vs torch
@@ -101,7 +99,7 @@ jobs:
- name: Run GPT2
run: |
BENCHMARK_LOG=gpt2_nojit JIT=0 python3.11 examples/gpt2.py --prompt "Hello." --count 10 --temperature 0 --timing | tee gpt2_unjitted.txt
BENCHMARK_LOG=gpt2 JIT=1 ASSERT_MIN_STEP_TIME=13 python3.11 examples/gpt2.py --prompt "Hello." --count 10 --temperature 0 --timing | tee gpt2_jitted.txt
BENCHMARK_LOG=gpt2 JIT=1 ASSERT_MIN_STEP_TIME=8 python3.11 examples/gpt2.py --prompt "Hello." --count 10 --temperature 0 --timing | tee gpt2_jitted.txt
- name: Run GPT2 w HALF
run: BENCHMARK_LOG=gpt2_half HALF=1 python3.11 examples/gpt2.py --count 10 --temperature 0 --timing | tee gpt2_half.txt
- name: Run GPT2 w HALF/BEAM
@@ -110,19 +108,14 @@ jobs:
run: BENCHMARK_LOG=olmoe python3.11 examples/olmoe.py
- name: Train MNIST
run: time PYTHONPATH=. TARGET_EVAL_ACC_PCT=96.0 python3.11 examples/beautiful_mnist.py | tee beautiful_mnist.txt
# NOTE: this is failing in CI. it is not failing on my machine and I don't really have a way to debug it
# the error is "RuntimeError: Internal Error (0000000e:Internal Error)"
#- name: Run 10 CIFAR training steps
# run: BENCHMARK_LOG=cifar_10steps JIT=1 ASSERT_MIN_STEP_TIME=3000 STEPS=10 python3.11 examples/hlb_cifar10.py | tee train_cifar.txt
#- name: Run 10 CIFAR training steps w HALF
# run: BENCHMARK_LOG=cifar_10steps_half JIT=2 ASSERT_MIN_STEP_TIME=3000 STEPS=10 DEFAULT_FLOAT=HALF python3.11 examples/hlb_cifar10.py | tee train_cifar_half.txt
- name: Run 10 CIFAR training steps
run: BENCHMARK_LOG=cifar_10steps JIT=1 ASSERT_MIN_STEP_TIME=320 STEPS=10 python3.11 examples/hlb_cifar10.py | tee train_cifar.txt
- name: Run 10 CIFAR training steps w HALF
run: BENCHMARK_LOG=cifar_10steps_half JIT=2 ASSERT_MIN_STEP_TIME=385 STEPS=10 DEFAULT_FLOAT=HALF python3.11 examples/hlb_cifar10.py | tee train_cifar_half.txt
#- name: Run 10 CIFAR training steps w BF16
# run: STEPS=10 DEFAULT_FLOAT=BFLOAT16 python3.11 examples/hlb_cifar10.py | tee train_cifar_bf16.txt
# TODO: too slow
# - name: Run 10 CIFAR training steps w winograd
# run: BENCHMARK_LOG=cifar_10steps_wino JIT=1 ASSERT_MIN_STEP_TIME=150 WINO=1 STEPS=10 python3.11 examples/hlb_cifar10.py | tee train_cifar_wino.txt
- name: Run 10 CIFAR training steps w winograd
run: BENCHMARK_LOG=cifar_10steps_wino JIT=1 ASSERT_MIN_STEP_TIME=150 WINO=1 STEPS=10 python3.11 examples/hlb_cifar10.py | tee train_cifar_wino.txt
- name: UsbGPU boot time
run: sudo -E PYTHONPATH=. DEBUG=2 AM_RESET=1 AMD=1 AMD_IFACE=USB time python3.11 test/test_tiny.py TestTiny.test_plus
- name: UsbGPU tiny tests
@@ -167,7 +160,7 @@ jobs:
testnvidiabenchmark:
name: tinybox green Benchmark
runs-on: [self-hosted, Linux, tinyboxgreen]
timeout-minutes: 60
timeout-minutes: 30
defaults:
run:
shell: bash -e -o pipefail {0}
@@ -220,9 +213,8 @@ jobs:
run: DEBUG=2 CUDA=1 python -m pytest -rA test/test_tiny.py
- name: Run Stable Diffusion
run: BENCHMARK_LOG=stable_diffusion NV=1 python3 examples/stable_diffusion.py --fp16 --seed 0 --noshow --timing | tee sd.txt
# TODO: too slow
# - name: Run SDXL
# run: BENCHMARK_LOG=stable_diffusion_xl ASSERT_MIN_STEP_TIME=2000 CAPTURE_PROCESS_REPLAY=0 NV=1 CAPTURE_PROCESS_REPLAY=0 python3 examples/sdxl.py --seed 0 --noshow --timing | tee sdxl.txt
- name: Run SDXL
run: BENCHMARK_LOG=stable_diffusion_xl ASSERT_MIN_STEP_TIME=2000 CAPTURE_PROCESS_REPLAY=0 NV=1 CAPTURE_PROCESS_REPLAY=0 python3 examples/sdxl.py --seed 0 --noshow --timing | tee sdxl.txt
- name: Run LLaMA
run: |
BENCHMARK_LOG=llama_nojit NV=1 JIT=0 python3 examples/llama.py --gen 1 --prompt "Hello." --count 10 --temperature 0 --timing | tee llama_unjitted.txt
@@ -246,9 +238,9 @@ jobs:
- name: Run GPT2
run: |
BENCHMARK_LOG=gpt2_nojit NV=1 JIT=0 python3 examples/gpt2.py --prompt "Hello." --count 10 --temperature 0 --timing | tee gpt2_unjitted.txt
BENCHMARK_LOG=gpt2 NV=1 JIT=1 ASSERT_MIN_STEP_TIME=4 python3 examples/gpt2.py --prompt "Hello." --count 10 --temperature 0 --timing | tee gpt2_jitted.txt
BENCHMARK_LOG=gpt2 NV=1 JIT=1 ASSERT_MIN_STEP_TIME=5 python3 examples/gpt2.py --prompt "Hello." --count 10 --temperature 0 --timing | tee gpt2_jitted.txt
- name: Run GPT2 w HALF
run: BENCHMARK_LOG=gpt2_half NV=1 HALF=1 ASSERT_MIN_STEP_TIME=6 python3 examples/gpt2.py --count 10 --temperature 0 --timing | tee gpt2_half.txt
run: BENCHMARK_LOG=gpt2_half NV=1 HALF=1 ASSERT_MIN_STEP_TIME=5 python3 examples/gpt2.py --count 10 --temperature 0 --timing | tee gpt2_half.txt
- name: Run GPT2 w HALF/BEAM
run: BENCHMARK_LOG=gpt2_half_beam NV=1 HALF=1 JITBEAM=2 IGNORE_BEAM_CACHE=1 python3 examples/gpt2.py --count 10 --temperature 0 --timing | tee gpt2_half_beam.txt
- uses: actions/upload-artifact@v4
@@ -282,7 +274,7 @@ jobs:
testmorenvidiabenchmark:
name: tinybox green Training Benchmark
runs-on: [self-hosted, Linux, tinyboxgreen]
timeout-minutes: 60
timeout-minutes: 20
defaults:
run:
shell: bash -e -o pipefail {0}
@@ -307,27 +299,24 @@ jobs:
rm -f /tmp/staging.db /tmp/staging.db-shm /tmp/staging.db-wal
- name: reset process replay
run: test/external/process_replay/reset.py
# TODO: too slow
# - name: Fuzz Padded Tensor Core GEMM (NV)
# run: NV=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
# TODO: too slow
# - name: Fuzz Padded Tensor Core GEMM (PTX)
# 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: Fuzz Padded Tensor Core GEMM (NV)
run: NV=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: Fuzz Padded Tensor Core GEMM (PTX)
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
run: time PYTHONPATH=. NV=1 TARGET_EVAL_ACC_PCT=96.0 python3 examples/beautiful_mnist.py | tee beautiful_mnist.txt
- 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=85 NV=1 STEPS=10 python3 examples/hlb_cifar10.py | tee train_cifar.txt
- name: Run 10 CIFAR training steps w HALF
run: BENCHMARK_LOG=cifar_10steps_half ASSERT_MIN_STEP_TIME=310 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=68 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
run: BENCHMARK_LOG=cifar_10steps_bf16 ASSERT_MIN_STEP_TIME=310 NV=1 STEPS=10 DEFAULT_FLOAT=BFLOAT16 python3 examples/hlb_cifar10.py | tee train_cifar_bf16.txt
# TODO: too slow
# - 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_bf16 ASSERT_MIN_STEP_TIME=75 NV=1 STEPS=10 DEFAULT_FLOAT=BFLOAT16 python3 examples/hlb_cifar10.py | tee train_cifar_bf16.txt
- name: Run 10 CIFAR training steps w winograd
run: BENCHMARK_LOG=cifar_10steps_half_wino ASSERT_MIN_STEP_TIME=35 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
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 LATEWINO=1 STEPS=1000 TARGET_EVAL_ACC_PCT=93.2 python3 examples/hlb_cifar10.py | tee train_cifar_one_gpu.txt
- 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.2 python3 examples/hlb_cifar10.py | tee train_cifar_six_gpu.txt
- name: Run MLPerf resnet eval on training data
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)
@@ -357,7 +346,7 @@ jobs:
testamdbenchmark:
name: tinybox red Benchmark
runs-on: [self-hosted, Linux, tinybox]
timeout-minutes: 60
timeout-minutes: 20
defaults:
run:
shell: bash -e -o pipefail {0}
@@ -426,10 +415,9 @@ jobs:
- name: Test AM warm start time
run: time AMD=1 python3 test/test_tiny.py TestTiny.test_plus
- name: Run Stable Diffusion
run: BENCHMARK_LOG=stable_diffusion ASSERT_MIN_STEP_TIME=550 AMD=1 python3 examples/stable_diffusion.py --fp16 --seed 0 --noshow --timing | tee sd.txt
# TODO: too slow
# - name: Run SDXL
# run: BENCHMARK_LOG=stable_diffusion_xl ASSERT_MIN_STEP_TIME=3200 CAPTURE_PROCESS_REPLAY=0 AMD=1 python3 examples/sdxl.py --seed 0 --noshow --timing | tee sdxl.txt
run: BENCHMARK_LOG=stable_diffusion ASSERT_MIN_STEP_TIME=450 AMD=1 python3 examples/stable_diffusion.py --fp16 --seed 0 --noshow --timing | tee sd.txt
- name: Run SDXL
run: BENCHMARK_LOG=stable_diffusion_xl ASSERT_MIN_STEP_TIME=1400 CAPTURE_PROCESS_REPLAY=0 AMD=1 python3 examples/sdxl.py --seed 0 --noshow --timing | tee sdxl.txt
- name: Run LLaMA 7B
run: |
BENCHMARK_LOG=llama_nojit AMD=1 JIT=0 python3 examples/llama.py --gen 1 --prompt "Hello." --count 10 --temperature 0 --timing | tee llama_unjitted.txt
@@ -488,7 +476,7 @@ jobs:
testmoreamdbenchmark:
name: tinybox red Training Benchmark
runs-on: [self-hosted, Linux, tinybox]
timeout-minutes: 60
timeout-minutes: 30
defaults:
run:
shell: bash -e -o pipefail {0}
@@ -520,20 +508,19 @@ jobs:
- name: Train MNIST
run: time PYTHONPATH=. AMD=1 TARGET_EVAL_ACC_PCT=96.0 python3 examples/beautiful_mnist.py | tee beautiful_mnist.txt
- 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=85 AMD=1 STEPS=10 python3 examples/hlb_cifar10.py | tee train_cifar.txt
- name: Run 10 CIFAR training steps w HALF
run: BENCHMARK_LOG=cifar_10steps_half ASSERT_MIN_STEP_TIME=330 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
# 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
# - 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 ASSERT_MIN_STEP_TIME=188 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
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
- 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
- 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 LATEWINO=1 STEPS=1000 TARGET_EVAL_ACC_PCT=93.2 python3 examples/hlb_cifar10.py | tee train_cifar_one_gpu.txt
#- 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.2 python3 examples/hlb_cifar10.py | tee train_cifar_six_gpu.txt
#- name: Run full CIFAR training steps w 6 GPUS (REMOTE)
# run: time BENCHMARK_LOG=cifar_6gpu_remote REMOTE=1 REMOTEDEV=AMD 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_remote.txt
# run: time BENCHMARK_LOG=cifar_6gpu_remote REMOTE=1 REMOTEDEV=AMD DEFAULT_FLOAT=HALF STEPS=350 BS=1536 GPUS=6 TARGET_EVAL_ACC_PCT=93.2 python3 examples/hlb_cifar10.py | tee train_cifar_six_gpu_remote.txt
- uses: actions/upload-artifact@v4
with:
name: Speed (AMD Training)
@@ -552,7 +539,7 @@ jobs:
testmlperfamdbenchmark:
name: tinybox red MLPerf Benchmark
runs-on: [self-hosted, Linux, tinybox]
timeout-minutes: 60
timeout-minutes: 30
defaults:
run:
shell: bash -e -o pipefail {0}
@@ -619,21 +606,21 @@ jobs:
- name: reset process replay
run: test/external/process_replay/reset.py
- name: benchmark openpilot 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 test/external/external_benchmark_openpilot.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/driving_vision.onnx
run: BENCHMARK_LOG=openpilot_0_9_9_vision ASSERT_MIN_STEP_TIME=30 PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 test/external/external_benchmark_openpilot.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/driving_vision.onnx
- name: benchmark openpilot 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 test/external/external_benchmark_openpilot.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/driving_policy.onnx
run: BENCHMARK_LOG=openpilot_0_9_9_policy ASSERT_MIN_STEP_TIME=45 PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 test/external/external_benchmark_openpilot.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/driving_policy.onnx
- name: benchmark openpilot 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 test/external/external_benchmark_openpilot.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/dmonitoring_model.onnx
run: BENCHMARK_LOG=openpilot_0_9_9_dmonitoring ASSERT_MIN_STEP_TIME=70 PYTHONPATH=. NOLOCALS=1 FLOAT16=1 IMAGE=2 QCOM=1 taskset -c 4-7 python3 test/external/external_benchmark_openpilot.py https://github.com/commaai/openpilot/raw/v0.9.9/selfdrive/modeld/models/dmonitoring_model.onnx
- name: openpilot compile3 0.9.9 driving_vision
run: PYTHONPATH="." ASSERT_MIN_STEP_TIME=22 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: PYTHONPATH="." 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
run: PYTHONPATH="." ASSERT_MIN_STEP_TIME=7 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: PYTHONPATH="." 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
run: PYTHONPATH="." ASSERT_MIN_STEP_TIME=15 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: PYTHONPATH="." 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 Space Lab policy + vision
run: |
PYTHONPATH="." ASSERT_MIN_STEP_TIME=4 QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://gitlab.com/commaai/openpilot-lfs.git/gitlab-lfs/objects/22aec22a10ce09384d4a4af2a0bbff08d54af7e0c888503508f356fae4ff0e29
PYTHONPATH="." ASSERT_MIN_STEP_TIME=26 QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://gitlab.com/commaai/openpilot-lfs.git/gitlab-lfs/objects/c824f68646a3b94f117f01c70dc8316fb466e05fbd42ccdba440b8a8dc86914b
PYTHONPATH="." QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://gitlab.com/commaai/openpilot-lfs.git/gitlab-lfs/objects/22aec22a10ce09384d4a4af2a0bbff08d54af7e0c888503508f356fae4ff0e29
PYTHONPATH="." QCOM=1 taskset -c 4-7 python3 examples/openpilot/compile3.py https://gitlab.com/commaai/openpilot-lfs.git/gitlab-lfs/objects/c824f68646a3b94f117f01c70dc8316fb466e05fbd42ccdba440b8a8dc86914b
- name: benchmark MobileNetV2 on DSP
run: |
# generate quantized weights
@@ -658,7 +645,7 @@ jobs:
testreddriverbenchmark:
name: AM Benchmark
runs-on: [self-hosted, Linux, tinyboxrandom]
timeout-minutes: 20
timeout-minutes: 15
defaults:
run:
shell: bash -e -o pipefail {0}
@@ -708,7 +695,7 @@ jobs:
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
- 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
run: time BENCHMARK_LOG=cifar AMD=1 DEFAULT_FLOAT=HALF LATEWINO=1 STEPS=1000 TARGET_EVAL_ACC_PCT=93.2 python3 examples/hlb_cifar10.py | tee am_train_cifar_one_gpu.txt
# TODO: enable
# - 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
@@ -729,7 +716,7 @@ jobs:
testgreendriverbenchmark:
name: NV Benchmark
runs-on: [self-hosted, Linux, tinyboxrandom]
timeout-minutes: 20
timeout-minutes: 15
defaults:
run:
shell: bash -e -o pipefail {0}
@@ -771,7 +758,7 @@ jobs:
- 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
- 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
run: time BENCHMARK_LOG=cifar NV=1 DEFAULT_FLOAT=HALF LATEWINO=1 STEPS=1000 TARGET_EVAL_ACC_PCT=93.2 python3 examples/hlb_cifar10.py | tee nv_train_cifar_one_gpu.txt
#- 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
- name: Run 10 MLPerf Bert training steps (1 gpu)
+54 -32
View File
@@ -30,6 +30,8 @@ jobs:
key: llvm-speed
deps: testing_minimal
llvm: 'true'
- name: External Benchmark Schedule
run: python3 test/external/external_benchmark_schedule.py
- name: Speed Test
run: CPU=1 CPU_LLVM=1 python3 test/speed/external_test_speed_v_torch.py
- name: Speed Test (BEAM=2)
@@ -46,7 +48,7 @@ jobs:
uses: ./.github/actions/setup-tinygrad
with:
deps: docs
pydeps: "capstone torch"
pydeps: "capstone"
- name: Build wheel and show size
run: |
pip install build
@@ -77,8 +79,6 @@ jobs:
run: |
python docs/abstractions2.py
python docs/abstractions3.py
- name: Test README
run: awk '/```python/{flag=1;next}/```/{flag=0}flag' README.md > README.py && python README.py
- name: Test Quickstart
run: awk '/```python/{flag=1;next}/```/{flag=0}flag' docs/quickstart.md > quickstart.py && python quickstart.py
- name: Test DEBUG
@@ -144,7 +144,7 @@ jobs:
sudo apt update || true
sudo apt install -y --no-install-recommends ninja-build
- name: Test beautiful_mnist in torch with TINY_BACKEND
run: CPU=1 CPU_LLVM=1 TARGET_EVAL_ACC_PCT=96.0 TINY_BACKEND=1 python3 examples/other_mnist/beautiful_mnist_torch.py
run: SPLIT_REDUCEOP=0 FUSE_ARANGE=1 CPU=1 CPU_LLVM=1 TARGET_EVAL_ACC_PCT=96.0 TINY_BACKEND=1 python3 examples/other_mnist/beautiful_mnist_torch.py
- name: Test some torch tests (expect failure)
run: python3 -m pytest extra/torch_backend/torch_tests.py -v --tb=no || true
@@ -259,28 +259,21 @@ jobs:
uses: ./.github/actions/setup-tinygrad
with:
key: unittest-12
pydeps: "pillow numpy ftfy regex"
pydeps: "pillow"
deps: testing_unit
- name: Check Device.DEFAULT
run: python -c "from tinygrad import Device; assert Device.DEFAULT == 'CPU', Device.DEFAULT"
- name: Test README
run: awk '/```python/{flag=1;next}/```/{flag=0}flag' README.md > README.py && python README.py
- name: Run unit tests
run: CPU=1 python -m pytest -n=auto test/unit/ --durations=20
- name: Check SPEC=1
run: SPEC=1 python3 test/test_tiny.py
run: python -m pytest -n=auto test/unit/ --durations=20
- name: Run targetted tests on NULL backend
run: NULL=1 python3 -m unittest test.test_multitensor.TestMultiTensor.test_data_parallel_resnet_train_step test/device/test_null.py
# TODO: too slow
# - name: Run SDXL on NULL backend
# run: NULL=1 DEBUG=1 python3 examples/sdxl.py --seed 0 --noshow --timing --fakeweights
- name: Run Clip tests for SD MLPerf on NULL backend
run: NULL=1 python -m pytest -n=auto test/external/mlperf_stable_diffusion/external_test_models.py::TestOpenClip --durations=20
run: NULL=1 python3 test/test_multitensor.py TestMultiTensor.test_data_parallel_resnet_train_step
- name: Run SDXL on NULL backend
run: MAX_BUFFER_SIZE=0 NULL=1 DEBUG=1 python3 examples/sdxl.py --seed 0 --noshow --timing --fakeweights
# TODO: support fake weights
#- name: Run LLaMA 7B on 4 fake devices
# run: NULL=1 python3 examples/llama.py --gen 1 --size 7B --shard 4 --prompt "Hello." --count 3 --temperature 0 --timing
- name: Run GC tests
run: python test/external/external_uop_gc.py
- name: External Benchmark Schedule
run: python3 test/external/external_benchmark_schedule.py
- name: Run process replay tests
uses: ./.github/actions/process-replay
- name: Regen dataset on test_tiny
@@ -317,9 +310,9 @@ jobs:
run: python test/external/fuzz_shape_ops.py
testopenclimage:
name: CL IMAGE Tests
name: 'CL IMAGE Tests'
runs-on: ubuntu-22.04
timeout-minutes: 15
timeout-minutes: 10
steps:
- name: Checkout Code
uses: actions/checkout@v4
@@ -337,7 +330,7 @@ jobs:
uses: ./.github/actions/process-replay
testgpumisc:
name: CL Misc tests
name: 'CL Misc tests'
runs-on: ubuntu-22.04
timeout-minutes: 10
steps:
@@ -362,7 +355,7 @@ jobs:
path: /tmp/sops.gz
testopenpilot:
name: openpilot Compile Tests
name: 'openpilot Compile Tests'
runs-on: ubuntu-22.04
timeout-minutes: 15
steps:
@@ -377,7 +370,7 @@ jobs:
llvm: 'true'
- name: Test openpilot model kernel count and gate usage
run: |
ALLOWED_KERNEL_COUNT=190 ALLOWED_READ_IMAGE=2041 ALLOWED_GATED_READ_IMAGE=41 FLOAT16=0 CL=1 IMAGE=2 python examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.9.4/selfdrive/modeld/models/supercombo.onnx
ALLOWED_KERNEL_COUNT=208 ALLOWED_READ_IMAGE=2175 ALLOWED_GATED_READ_IMAGE=16 FLOAT16=0 CL=1 IMAGE=2 python examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/v0.9.4/selfdrive/modeld/models/supercombo.onnx
- name: Test openpilot alt model correctness (float32)
run: FLOAT16=0 DEBUGCL=1 CL=1 IMAGE=2 python examples/openpilot/compile3.py https://github.com/commaai/openpilot/raw/3799fe46b3a629e491d4b8498b8ae83e4c88c304/selfdrive/modeld/models/supercombo.onnx
- name: Test openpilot fastvits model correctness (float32)
@@ -394,7 +387,7 @@ jobs:
# ****** ONNX Tests ******
testonnxcpu:
name: ONNX (CPU) Tests
name: 'ONNX (CPU) Tests'
runs-on: ubuntu-22.04
timeout-minutes: 20
@@ -422,7 +415,7 @@ jobs:
uses: ./.github/actions/process-replay
testopencl:
name: ONNX (CL)+Optimization Tests
name: 'ONNX (GPU)+Optimization Tests'
runs-on: ubuntu-22.04
timeout-minutes: 20
steps:
@@ -446,12 +439,8 @@ jobs:
run: CL=1 IGNORE_BEAM_CACHE=1 python3 -m pytest extra/optimization/test_beam_search.py
- name: Test MLPerf stuff
run: CL=1 python -m pytest -n=auto test/external/external_test_optim.py test/external/external_test_losses.py test/external/external_test_metrics.py test/external/external_test_datasets.py --durations=20
- name: NULL=1 beautiful_mnist_multigpu
run: NULL=1 python examples/beautiful_mnist_multigpu.py
- name: Test Bert training
run: NULL=1 DEFAULT_FLOAT=HALF BENCHMARK=10 BS=24 GPUS=4 BERT_LAYERS=2 MODEL=bert python3 examples/mlperf/model_train.py
- name: Test llama 3 training
run: NULL=1 SAMPLES=300 BS=8 SEQLEN=512 GRADIENT_ACC_STEPS=8 FAKEDATA=1 DEFAULT_FLOAT=bfloat16 OPTIM_DTYPE=bfloat16 LLAMA3_SIZE=1B MODEL=llama3 python3 examples/mlperf/model_train.py
run: MAX_BUFFER_SIZE=0 DEV=NULL SAMPLES=300 BS=8 SEQLEN=512 GRADIENT_ACC_STEPS=8 FAKEDATA=1 DEFAULT_FLOAT=bfloat16 OPTIM_DTYPE=bfloat16 LLAMA3_SIZE=1B MODEL=llama3 python3 examples/mlperf/model_train.py
- name: Run process replay tests
uses: ./.github/actions/process-replay
@@ -514,6 +503,39 @@ jobs:
# ****** Feature Tests ******
testrangeify:
name: Linux (rangeify)
runs-on: ubuntu-24.04
timeout-minutes: 15
steps:
- name: Checkout Code
uses: actions/checkout@v4
- name: Setup Environment
uses: ./.github/actions/setup-tinygrad
with:
key: rangeify-minimal-llvm
deps: testing_minimal
opencl: 'true'
llvm: "true"
- name: Test CPU=1 RANGEIFY=1
# TODO: add more passing tests here
# test_symbolic_arange_sym_step is passing now
# test_threefry_doesnt_use_long is because there's a contig after the long now
run: |
CPU=1 CPU_LLVM=0 RANGEIFY=1 python3 -m pytest -n auto --durations 20 \
-k "not test_symbolic_arange_sym_step and not test_threefry_doesnt_use_long" \
test/test_tiny.py test/test_rangeify.py test/test_ops.py test/test_tensor_variable.py \
test/test_outerworld_range.py test/test_sample.py test/test_randomness.py
- name: Test multitensor
run: RANGEIFY=1 PYTHONPATH="." python3 test/test_multitensor.py TestMultiTensor.test_matmul_shard_1_1 TestMultiTensor.test_simple_add_W
- name: Test GPU=1 RANGEIFY=1
run: GPU=1 RANGEIFY=1 pytest -n auto test/test_ops.py
- name: Test CPU=1 RANGEIFY=2
run: CPU=1 CPU_LLVM=0 RANGEIFY=2 python3 -m pytest -n auto test/test_tiny.py test/test_rangeify.py test/test_ops.py --durations 20
# slow (and still wrong on beautiful_mnist)
#- name: Test LLVM=1 RANGEIFY=1 (slow tests)
# run: CPU=1 CPU_LLVM=1 RANGEIFY=1 python3 -m pytest -n auto test/models/test_mnist.py --durations 20
testdevectorize:
name: Linux (devectorize)
runs-on: ubuntu-24.04
@@ -533,7 +555,7 @@ jobs:
- name: Test LLVM=1 DEVECTORIZE=0 for model
run: CPU=1 CPU_LLVM=1 DEVECTORIZE=0 python3 test/models/test_efficientnet.py
- name: Test CPU=1 DEVECTORIZE=0
run: CPU=1 CPU_LLVM=0 DEVECTORIZE=0 python3 -m pytest -n auto test/test_tiny.py test/test_ops.py -k "not test_avg_pool3d_failure"
run: CPU=1 CPU_LLVM=0 DEVECTORIZE=0 FUSE_ARANGE=0 python3 -m pytest -n auto test/test_tiny.py test/test_ops.py -k "not test_avg_pool3d_failure"
testdsp:
name: Linux (DSP)
@@ -634,7 +656,7 @@ jobs:
run: TRANSCENDENTAL=2 python -m pytest -n=auto test/test_ops.py::TestOps::test_sin test/test_ops.py::TestOps::test_cos test/test_ops.py::TestOps::test_tan test/test_ops.py::TestOps::test_exp test/test_ops.py::TestOps::test_log --durations=20
- name: Run TestOps.test_add with SQTT
run: |
VIZ=1 SQTT=1 DEBUG=5 python3 test/test_ops.py TestOps.test_add
PROFILE=1 SQTT=1 DEBUG=5 python3 test/test_ops.py TestOps.test_add
extra/sqtt/rgptool.py create "/tmp/profile.pkl.$USER" -o /tmp/gpu0.rgp
- name: Run process replay tests
uses: ./.github/actions/process-replay
+1 -20
View File
@@ -414,29 +414,10 @@ 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 extra/sqtt/rocprof/rocprof.py
fixup extra/sqtt/rocprof/rocprof.py
sed -i '1s/^/# pylint: skip-file\n/' extra/sqtt/rocprof/rocprof.py
sed -i "s/import ctypes/import ctypes, tinygrad.helpers.fetch as tgfetch/g" extra/sqtt/rocprof/rocprof.py
sed -i "s|FunctionFactoryStub()|ctypes.CDLL(str(tgfetch('https://github.com/ROCm/rocprof-trace-decoder/raw/5420409ad0963b2d76450add067b9058493ccbd0/releases/linux_glibc_2_28_x86_64/librocprof-trace-decoder.so')))|g" extra/sqtt/rocprof/rocprof.py
}
generate_webgpu() {
+7 -7
View File
@@ -42,6 +42,7 @@ import struct
from tinygrad.dtype import dtypes
from tinygrad.device import Buffer, Device
from tinygrad.uop.ops import UOp, Ops
from tinygrad.shape.shapetracker import ShapeTracker
# allocate some buffers + load in values
out = Buffer(DEVICE, 1, dtypes.int32).allocate()
@@ -50,14 +51,13 @@ b = Buffer(DEVICE, 1, dtypes.int32).allocate().copyin(memoryview(bytearray(struc
# NOTE: a._buf is the same as the return from cpu.allocator.alloc
# describe the computation
idx = UOp.const(dtypes.index, 0)
buf_1 = UOp(Ops.DEFINE_GLOBAL, dtypes.int32.ptr(), (), 1)
buf_2 = UOp(Ops.DEFINE_GLOBAL, dtypes.int32.ptr(), (), 2)
ld_1 = UOp(Ops.LOAD, dtypes.int32, (buf_1.index(idx),))
ld_2 = UOp(Ops.LOAD, dtypes.int32, (buf_2.index(idx),))
ld_1 = UOp(Ops.LOAD, dtypes.int32, (buf_1.view(ShapeTracker.from_shape((1,))),))
ld_2 = UOp(Ops.LOAD, dtypes.int32, (buf_2.view(ShapeTracker.from_shape((1,))),))
alu = ld_1 + ld_2
output_buf = UOp(Ops.DEFINE_GLOBAL, dtypes.int32.ptr(), (), 0)
st_0 = UOp(Ops.STORE, dtypes.void, (output_buf.index(idx), alu))
st_0 = UOp(Ops.STORE, dtypes.void, (output_buf.view(ShapeTracker.from_shape((1,))), alu))
s = UOp(Ops.SINK, dtypes.void, (st_0,))
# convert the computation to a "linearized" format (print the format)
@@ -80,7 +80,7 @@ print("******** third, the UOp ***********")
from tinygrad.engine.realize import run_schedule
from tinygrad.engine.schedule import create_schedule_with_vars
from tinygrad.schedule.rangeify import get_rangeify_map
from tinygrad.schedule.kernelize import get_kernelize_map
# allocate some values + load in values
a = UOp.new_buffer(DEVICE, 1, dtypes.int32)
@@ -93,10 +93,10 @@ out = a + b
s = UOp(Ops.SINK, dtypes.void, (out,))
# group the computation into kernels
becomes_map = get_rangeify_map(s)
becomes_map = get_kernelize_map(s)
# the compute maps to an assign
assign = becomes_map[a+b].base
assign = becomes_map[a+b]
# the first source is the output buffer (data)
assert assign.src[0].op is Ops.BUFFER
+1 -1
View File
@@ -10,7 +10,7 @@ Directories are listed in order of how they are processed.
Group UOps into kernels.
::: tinygrad.schedule.rangeify.get_rangeify_map
::: tinygrad.schedule.kernelize.get_kernelize_map
options:
members: false
show_labels: false
+2
View File
@@ -41,6 +41,8 @@ BEAM | [#] | number of beams in kernel beam search
DEFAULT_FLOAT | [HALF, ...]| specify the default float dtype (FLOAT32, HALF, BFLOAT16, FLOAT64, ...), default to FLOAT32
IMAGE | [1-2] | enable 2d specific optimizations
FLOAT16 | [1] | use float16 for images instead of float32
PTX | [1] | enable the specialized [PTX](https://docs.nvidia.com/cuda/parallel-thread-execution/) assembler for Nvidia GPUs. If not set, defaults to generic CUDA codegen backend.
PROFILE | [1] | enable profiling. This feature is supported in NV, AMD, QCOM and METAL backends.
VISIBLE_DEVICES | [list[int]]| restricts the NV/AMD devices that are available. The format is a comma-separated list of identifiers (indexing starts with 0).
JIT | [0-2] | 0=disabled, 1=[jit enabled](quickstart.md#jit) (default), 2=jit enabled, but graphs are disabled
VIZ | [1] | 0=disabled, 1=[viz enabled](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/viz)
+11 -18
View File
@@ -2,17 +2,17 @@
tinygrad supports various runtimes, enabling your code to scale across a wide range of devices. The default runtime can be automatically selected based on the available hardware, or you can force a specific runtime to be default using environment variables (e.g., `CPU=1`).
| Runtime | Description | Compiler Options | Requirements |
|---------|-------------|------------------|--------------|
| [NV](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_nv.py) | Provides acceleration for NVIDIA GPUs | nvrtc (default)<br>PTX (`NV_PTX=1`) | Ampere/Ada/Blackwell series GPUs.<br>You can select an interface via `NV_IFACE=(NVK\|PCI)`. See [NV interfaces](#nv-interfaces) for details. |
| [AMD](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_amd.py) | Provides acceleration for AMD GPUs | LLVM (`AMD_LLVM=1`)<br>HIP/COMGR (`AMD_HIP=1`) | RDNA2 or newer GPUs.<br>You can select an interface via `AMD_IFACE=(KFD\|PCI\|USB)`. See [AMD interfaces](#amd-interfaces) for details. |
| [QCOM](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_qcom.py) | Provides acceleration for QCOM GPUs | - | 6xx series GPUs |
| [METAL](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_metal.py) | Utilizes Metal for acceleration on Apple devices | - | M1+ Macs; Metal 3.0+ for `bfloat` support |
| [CUDA](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_cuda.py) | Utilizes CUDA for acceleration on NVIDIA GPUs | nvrtc (default)<br> PTX (`CUDA_PTX=1`) | NVIDIA GPU with CUDA support |
| [CL](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_cl.py) | Accelerates computations using OpenCL on GPUs | - | OpenCL 2.0 compatible device |
| [CPU](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_cpu.py) | Runs on CPU using the clang or llvm compiler | Clang JIT (default)<br>LLVM IR (`CPU_LLVM=1`) | `clang` compiler in system `PATH` |
| [WEBGPU](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_webgpu.py) | Runs on GPU using the Dawn WebGPU engine (used in Google Chrome) | - | Dawn library installed and discoverable. Binaries: [pydawn v0.3.0](https://github.com/wpmed92/pydawn/releases/tag/v0.3.0) |
| Runtime | Description | Requirements |
|---------|-------------|--------------|
| [NV](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_nv.py) | Provides acceleration for NVIDIA GPUs | Ampere/Ada series GPUs |
| [AMD](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_amd.py) | Provides acceleration for AMD GPUs | RDNA2/RDNA3/RDNA4 series GPUs. You can select one of the interfaces for communication by setting `AMD_IFACE=(KFD|PCI)`. See [AMD interfaces](#amd-interfaces) for more details. |
| [QCOM](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_qcom.py) | Provides acceleration for QCOM GPUs | 6xx series GPUs |
| [METAL](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_metal.py) | Utilizes Metal for acceleration on Apple devices | M1+ Macs; Metal 3.0+ for `bfloat` support |
| [CUDA](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_cuda.py) | Utilizes CUDA for acceleration on NVIDIA GPUs | NVIDIA GPU with CUDA support |
| [OpenCL](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_cl.py) | Accelerates computations using OpenCL on GPUs | OpenCL 2.0 compatible device |
| [CPU (C Code)](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_cpu.py) | Runs on CPU using the clang compiler | `clang` compiler in system `PATH` |
| [LLVM (LLVM IR)](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_llvm.py) | Runs on CPU using the LLVM compiler infrastructure | llvm libraries installed and findable |
| [WEBGPU](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/ops_webgpu.py) | Runs on GPU using the Dawn WebGPU engine (used in Google Chrome) | Dawn library installed and findable. Download binaries [here](https://github.com/wpmed92/pydawn/releases/tag/v0.3.0). |
## Interoperability
@@ -70,12 +70,5 @@ AMD backend supports several interfaces for communicating with devices:
* `KFD`: uses the amdgpu driver
* `PCI`: uses the [AM driver](developer/am.md)
* `USB`: USB3 interafce for asm24xx chips.
You can force an interface by setting `AMD_IFACE` to one of these values. In the case of `AMD_IFACE=PCI`, this may unbind your GPU from the amdgpu driver.
## NV Interfaces
NV backend supports several interfaces for communicating with devices:
* `NVK`: uses the nvidia driver
* `PCI`: uses the [NV driver](https://github.com/tinygrad/tinygrad/tree/master/tinygrad/runtime/support/nv/nvdev.py)
+1 -1
View File
@@ -10,7 +10,7 @@ GPUS = [f'{Device.DEFAULT}:{i}' for i in range(getenv("GPUS", 1))]
# override tinygrad defaults
dtypes.default_float = dtypes.half
Context(FUSE_OPTIM=1).__enter__()
Context(FUSE_ARANGE=1, FUSE_OPTIM=1).__enter__()
# from https://github.com/tysam-code/hlb-CIFAR10/blob/main/main.py
batchsize = getenv("BS", 1024)
+1 -1
View File
@@ -1,6 +1,6 @@
import sys, time
from tinygrad import TinyJit, GlobalCounters, fetch, getenv
from tinygrad.nn.onnx import OnnxRunner
from tinygrad.frontend.onnx import OnnxRunner
from extra.onnx_helpers import get_example_inputs, validate
def load_onnx_model(onnx_file):
+1 -1
View File
@@ -8,7 +8,7 @@ import numpy as np
import subprocess
import tensorflow as tf
import tf2onnx
from tinygrad.nn.onnx import OnnxRunner
from tinygrad.frontend.onnx import OnnxRunner
from tinygrad.tensor import Tensor
from tinygrad.helpers import to_mv
from extra.export_model import export_model_clang, compile_net, jit_model
+6 -6
View File
@@ -26,8 +26,8 @@ class Attention:
start_pos = start_pos.val
if HALF: x = x.half()
xqkv = self.c_attn(x).reshape(None, None, 3, self.n_heads, self.head_dim)
xq, xk, xv = [xqkv[:, :, i, :, :] for i in range(3)]
xqkv = self.c_attn(x)
xq, xk, xv = [xqkv.shrink((None, None, (i*self.dim, (i+1)*self.dim))).reshape(None, None, self.n_heads, self.head_dim) for i in range(3)]
bsz, seqlen, _, _ = xq.shape
# create kv cache
@@ -35,11 +35,11 @@ class Attention:
self.cache_kv = Tensor.zeros(2, bsz, MAX_CONTEXT, self.n_heads, self.head_dim, dtype=x.dtype).contiguous().realize()
# update the cache
self.cache_kv[:, :, start_pos:start_pos+seqlen, :, :].assign(Tensor.stack(xk, xv)).realize()
self.cache_kv.shrink((None, None,(start_pos,start_pos+seqlen),None,None)).assign(Tensor.stack(xk, xv)).realize()
if start_pos > 0:
keys = self.cache_kv[0][:, :start_pos+seqlen, :, :]
values = self.cache_kv[1][:, :start_pos+seqlen, :, :]
keys = self.cache_kv[0].shrink((None, (0, start_pos+seqlen), None, None))
values = self.cache_kv[1].shrink((None, (0, start_pos+seqlen), None, None))
else:
keys = xk
values = xv
@@ -64,7 +64,7 @@ class TransformerBlock:
def __call__(self, x:Tensor, start_pos:Variable, mask:Optional[Tensor]):
h = x + self.attn(self.ln_1(x), start_pos, mask).float()
return (h + self.mlp(self.ln_2(h))).contiguous()
return (h + self.mlp(self.ln_2(h)))
class Transformer:
def __init__(self, dim, n_heads, n_layers, norm_eps, vocab_size, max_seq_len=1024):
+2 -2
View File
@@ -145,6 +145,7 @@ hyp = {
},
}
@Context(FUSE_ARANGE=getenv("FUSE_ARANGE", 1))
def train_cifar():
def set_seed(seed):
@@ -228,8 +229,7 @@ def train_cifar():
if getenv("RANDOM_CROP", 1):
X = random_crop(X, crop_size=32)
if getenv("RANDOM_FLIP", 1):
# NOTE: RANGEIFY=1 needs this contiguous or the X[perms] is very slow
X = (Tensor.rand(X.shape[0],1,1,1) < 0.5).where(X.flip(-1), X).contiguous() # flip LR
X = (Tensor.rand(X.shape[0],1,1,1) < 0.5).where(X.flip(-1), X) # flip LR
X, Y = X[perms], Y[perms]
return X, Y, *cutmix(X, Y, perms, mask_size=hyp['net']['cutmix_size'])
-27
View File
@@ -511,33 +511,6 @@ def batch_load_retinanet(dataset, val:bool, base_dir:Path, batch_size:int=32, sh
# happens with BENCHMARK set
pass
# stable diffusion callbacks to match mlperf ref; declared here because they're pickled
def filter_dataset(sample:dict): return {k:v for k,v in sample.items() if k in {'npy', 'txt'}}
def collate(batch:list[dict]):
ret = {"npy": [], "txt": [], "__key__": []}
for sample in batch:
for k,v in sample.items():
ret[k].append(v)
return ret
def collate_fn(batch): return batch
# Reference (code): https://github.com/mlcommons/training/blob/2f4a93fb4888180755a8ef55f4b977ef8f60a89e/stable_diffusion/ldm/data/webdatasets.py, Line 55
# Reference (params): https://github.com/mlcommons/training/blob/ab4ae1ca718d7fe62c369710a316dff18768d04b/stable_diffusion/configs/train_01x08x08.yaml, Line 107
def batch_load_train_stable_diffusion(urls:str, BS:int):
import webdataset
dataset = webdataset.WebDataset(urls=urls, resampled=True, cache_size=-1, cache_dir=None)
dataset = dataset.shuffle(size=1000)
dataset = dataset.decode()
dataset = dataset.map(filter_dataset)
dataset = dataset.batched(BS, partial=False, collation_fn=collate)
dataset = webdataset.WebLoader(dataset, batch_size=None, shuffle=False, num_workers=1, persistent_workers=True, collate_fn=collate_fn)
for x in dataset:
assert isinstance(x, dict) and all(isinstance(k, str) for k in x.keys()) and all(isinstance(v, list) for v in x.values())
assert all(isinstance(moment_mean_logvar, np.ndarray) and moment_mean_logvar.shape==(1,8,64,64) for moment_mean_logvar in x["npy"])
assert all(isinstance(caption, str) for caption in x["txt"])
yield x
# llama3
class BinIdxDataset:
+1 -63
View File
@@ -2,9 +2,7 @@ import math
from typing import Union
from tinygrad import Tensor, nn, dtypes
from tinygrad.helpers import prod, argfix, Context
from tinygrad.nn.state import get_parameters
from extra.models.unet import UNetModel
from tinygrad.helpers import prod, argfix
# rejection sampling truncated randn
def rand_truncn(*shape, dtype=None, truncstds=2, **kwargs) -> Tensor:
@@ -19,10 +17,6 @@ def he_normal(*shape, a: float = 0.00, **kwargs) -> Tensor:
std = math.sqrt(2.0 / (1 + a ** 2)) / math.sqrt(prod(argfix(*shape)[1:])) / 0.87962566103423978
return std * rand_truncn(*shape, **kwargs)
# Stable Diffusion v2 training uses default torch gelu, which doesn't use tanh approximation
def gelu_erf(x:Tensor) -> Tensor:
return 0.5 * x * (1.0 + (x / 1.4142135623730951).erf())
class Conv2dHeNormal(nn.Conv2d):
def __init__(self, in_channels, out_channels, kernel_size, stride=1, padding=0, dilation=1, groups=1, bias=True):
super().__init__(in_channels, out_channels, kernel_size, stride=stride, padding=padding, dilation=dilation, groups=groups, bias=bias)
@@ -133,59 +127,3 @@ class Conv2dRetinaNet(nn.Conv2d):
def __call__(self, x:Tensor) -> Tensor:
return x.conv2d(self.weight.cast(dtypes.default_float), self.bias.cast(dtypes.default_float) if self.bias is not None else None,
groups=self.groups, stride=self.stride, dilation=self.dilation, padding=self.padding)
# copy torch AMP: isolate mixed precision to just the below autocast ops, instead of using dtypes.default_float which affects all new Tensors
class AutocastLinear(nn.Linear):
cast_dtype=dtypes.bfloat16 # enable monkeypatching of the mixed precision dtype
def __call__(self, x:Tensor) -> Tensor:
dtype = type(self).cast_dtype
return x.cast(dtype).linear(self.weight.cast(dtype).transpose(), self.bias.cast(dtype) if self.bias is not None else None)
class AutocastConv2d(nn.Conv2d):
cast_dtype=dtypes.bfloat16
def __call__(self, x:Tensor) -> Tensor:
dtype = type(self).cast_dtype
return x.cast(dtype).conv2d(self.weight.cast(dtype), self.bias.cast(dtype), self.groups, self.stride, self.dilation, self.padding)
# copy torch AMP: upcast to float32 before GroupNorm and LayerNorm
class AutocastGroupNorm(nn.GroupNorm):
def __call__(self, x:Tensor) -> Tensor:
return super().__call__(x.cast(dtypes.float32))
class AutocastLayerNorm(nn.LayerNorm):
def __call__(self, x:Tensor) -> Tensor:
return super().__call__(x.cast(dtypes.float32))
def zero_module(module):
for p in get_parameters(module): p.assign(Tensor.zeros_like(p).contiguous())
# Stable Diffusion mlperf reference doesn't call scaled_dot_product_attention
# copy torch AMP: upcast to float32 before softmax on CUDA
def attn_f32_softmax(q:Tensor, k:Tensor, v:Tensor) -> Tensor:
return (q.matmul(k.transpose(-2,-1), dtype=dtypes.float32) / math.sqrt(q.shape[-1])).softmax(-1).cast(q.dtype) @ v
def init_stable_diffusion(version:str, pretrained:str, devices:list[str]):
from examples.stable_diffusion import StableDiffusion
from tinygrad.nn.state import safe_load, safe_save, load_state_dict, get_state_dict
from tempfile import TemporaryDirectory
model = StableDiffusion(version=version, pretrained=pretrained)
unet:UNetModel = model.model.diffusion_model
# this prevents extra consumption of memory, enabling much larger BS
Tensor.realize(*get_parameters(unet))
with TemporaryDirectory(prefix="unet_init") as tmp:
safe_save(get_state_dict(unet), init_fn:=f"{tmp}/init_model.safetensors")
load_state_dict(unet, safe_load(init_fn))
sqrt_alphas_cumprod = model.alphas_cumprod.sqrt().realize()
sqrt_one_minus_alphas_cumprod = (1 - model.alphas_cumprod).sqrt().realize()
if len(devices) > 1:
to_move = [sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod]
if version == "v2-mlperf-train": to_move += get_parameters(unet) + get_parameters(model.cond_stage_model)
for p in to_move:
p.to_(devices)
with Context(BEAM=0):
Tensor.realize(*to_move)
return model, unet, sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod
+2 -23
View File
@@ -1,9 +1,8 @@
import math
from tinygrad import dtypes, Tensor
from tinygrad import dtypes
from tinygrad.nn.optim import Optimizer
from extra.lr_scheduler import LR_Scheduler
from typing import Callable
# https://github.com/mlcommons/training/blob/e237206991d10449d9675d95606459a3cb6c21ad/image_classification/tensorflow2/lars_util.py
class PolynomialDecayWithWarmup(LR_Scheduler):
@@ -37,24 +36,4 @@ class CosineAnnealingLRWithWarmup(LR_Scheduler):
def get_lr(self):
warmup_lr = ((self.epoch_counter+1) / self.warmup_steps) * self.base_lr
decay_lr = self.end_lr + 0.5 * (self.base_lr-self.end_lr) * (1 + (((self.epoch_counter+1-self.warmup_steps)/self.decay_steps) * math.pi).cos())
return (self.epoch_counter < self.warmup_steps).where(warmup_lr, decay_lr).cast(self.optimizer.lr.dtype)
# Reference: https://github.com/mlcommons/training/blob/64b14a9abc74e08779a175abca7d291f8c957632/stable_diffusion/ldm/lr_scheduler.py, Lines 36-97
class LambdaLinearScheduler:
def __init__(self, warm_up_steps:int, f_min:float, f_max:float, f_start:float, cycle_lengths:int):
self.lr_warm_up_steps, self.f_min, self.f_max, self.f_start, self.cycle_lengths = warm_up_steps, f_min, f_max, f_start, cycle_lengths
def schedule(self, n:Tensor) -> Tensor:
warm_up = (n < self.lr_warm_up_steps)
f_warm_up = (self.f_max - self.f_start) / self.lr_warm_up_steps * n + self.f_start
return warm_up.where(f_warm_up, self.f_min + (self.f_max - self.f_min) * (self.cycle_lengths - n) / (self.cycle_lengths))
# based on torch.optim.lr_scheduler.LambdaLR
class LambdaLR(LR_Scheduler):
def __init__(self, optimizer:Optimizer, base_lr:Tensor, lr_lambda:Callable):
super().__init__(optimizer)
self.base_lr, self.lr_lambda = base_lr, lr_lambda
self.step()
def get_lr(self):
return self.base_lr * self.lr_lambda(self.epoch_counter - 1)
return (self.epoch_counter < self.warmup_steps).where(warmup_lr, decay_lr).cast(self.optimizer.lr.dtype)
+2 -252
View File
@@ -1,10 +1,10 @@
import time, math, os
import time, math
start = time.perf_counter()
from pathlib import Path
import numpy as np
from tinygrad import Tensor, Device, dtypes, GlobalCounters, TinyJit
from tinygrad.nn.state import get_parameters, load_state_dict, safe_load
from tinygrad.helpers import getenv, Context, prod
from tinygrad.helpers import getenv
from extra.bench_log import BenchEvent, WallTimeEvent
def tlog(x): print(f"{x:25s} @ {time.perf_counter()-start:5.2f}s")
@@ -287,256 +287,6 @@ def eval_llama3():
log_perplexity = np.mean(losses)
print(f"Log Perplexity: {log_perplexity}")
# NOTE: BEAM hangs on 8xmi300x with DECODE_BS=384 in final realize below; function is declared here for external testing
@TinyJit
def vae_decode(x:Tensor, vae, disable_beam=False) -> Tensor:
from examples.stable_diffusion import AutoencoderKL
assert isinstance(vae, AutoencoderKL)
x = vae.post_quant_conv(1./0.18215 * x)
x = vae.decoder.conv_in(x)
x = vae.decoder.mid(x)
for i, l in enumerate(vae.decoder.up[::-1]):
print("decode", x.shape)
for b in l['block']: x = b(x)
if 'upsample' in l:
bs,c,py,px = x.shape
x = x.reshape(bs, c, py, 1, px, 1).expand(bs, c, py, 2, px, 2).reshape(bs, c, py*2, px*2)
x = l['upsample']['conv'](x)
if i == len(vae.decoder.up) - 1 and disable_beam:
with Context(BEAM=0): x.realize()
else: x.realize()
x = vae.decoder.conv_out(vae.decoder.norm_out(x).swish())
x = ((x + 1.0) / 2.0).clip(0.0, 1.0)
return x
def eval_stable_diffusion():
import csv, PIL, sys
from tqdm import tqdm
from examples.mlperf.initializers import init_stable_diffusion, gelu_erf
from examples.stable_diffusion import AutoencoderKL
from extra.models.unet import UNetModel
from tinygrad.nn.state import load_state_dict, torch_load
from tinygrad.helpers import BEAM
from extra.models import clip
from extra.models.clip import FrozenOpenClipEmbedder
from extra.models.clip import OpenClipEncoder
from extra.models.inception import FidInceptionV3
config = {}
GPUS = config["GPUS"] = [f"{Device.DEFAULT}:{i}" for i in range(getenv("GPUS", 1))]
for x in GPUS: Device[x]
print(f"running eval on {GPUS}")
seed = config["seed"] = getenv("SEED", 12345)
CKPTDIR = config["CKPTDIR"] = Path(getenv("CKPTDIR", "./checkpoints"))
DATADIR = config["DATADIR"] = Path(getenv("DATADIR", "./datasets"))
CONTEXT_BS = config["CONTEXT_BS"] = getenv("CONTEXT_BS", 1 * len(GPUS))
DENOISE_BS = config["DENOISE_BS"] = getenv("DENOISE_BS", 1 * len(GPUS))
DECODE_BS = config["DECODE_BS"] = getenv("DECODE_BS", 1 * len(GPUS))
INCEPTION_BS = config["INCEPTION_BS"] = getenv("INCEPTION_BS", 1 * len(GPUS))
CLIP_BS = config["CLIP_BS"] = getenv("CLIP_BS", 1 * len(GPUS))
EVAL_CKPT_DIR = config["EVAL_CKPT_DIR"] = getenv("EVAL_CKPT_DIR", "")
STOP_IF_CONVERGED = config["STOP_IF_CONVERGED"] = getenv("STOP_IF_CONVERGED", 0)
if (WANDB := getenv("WANDB", "")):
import wandb
wandb.init(config=config, project="MLPerf-Stable-Diffusion")
assert EVAL_CKPT_DIR != "", "provide a directory with checkpoints to be evaluated"
print(f"running eval on checkpoints in {EVAL_CKPT_DIR}\nSEED={seed}")
eval_queue:list[tuple[int, Path]] = []
for p in Path(EVAL_CKPT_DIR).iterdir():
if p.name.endswith(".safetensors"):
ckpt_iteration = p.name.split(".safetensors")[0]
assert ckpt_iteration.isdigit(), f"invalid checkpoint name: {p.name}, expected <digits>.safetensors"
eval_queue.append((int(ckpt_iteration), p))
assert len(eval_queue), f'no files ending with ".safetensors" were found in {EVAL_CKPT_DIR}'
print(sorted(eval_queue, reverse=True))
Tensor.manual_seed(seed) # seed for weight initialization
model, unet, sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod = init_stable_diffusion("v2-mlperf-eval", CKPTDIR / "sd" / "512-base-ema.ckpt", GPUS)
# load prompts for generating images for validation; 2 MB of data total
with open(DATADIR / "coco2014" / "val2014_30k.tsv") as f:
reader = csv.DictReader(f, delimiter="\t")
eval_inputs:list[dict] = [{"image_id": int(row["image_id"]), "id": int(row["id"]), "caption": row["caption"]} for row in reader]
assert len(eval_inputs) == 30_000
# NOTE: the clip weights are the same between model.cond_stage_model and clip_encoder
eval_timesteps = list(reversed(range(1, 1000, 20)))
original_device, Device.DEFAULT = Device.DEFAULT, "CPU"
# The choice of alphas_prev[0] = alphas_cumprod[0] seems arbitrary, but it's how the mlperf ref does it:
# alphas_prev = np.asarray([alphacums[0]] + alphacums[ddim_timesteps[:-1]].tolist())
eval_alphas_prev = model.alphas_cumprod[0:1].cat(model.alphas_cumprod[list(range(1, 1000, 20))[:-1]]).to(GPUS).realize()
inception = FidInceptionV3().load_from_pretrained(CKPTDIR / "inception" / "pt_inception-2015-12-05-6726825d.pth")
vision_cfg = {'width': 1280, 'layers': 32, 'd_head': 80, 'image_size': 224, 'patch_size': 14}
text_cfg = {'width': 1024, 'n_heads': 16, 'layers': 24, 'vocab_size': 49408, 'ctx_length': 77}
clip.gelu = gelu_erf
clip_encoder = OpenClipEncoder(1024, text_cfg, vision_cfg)
loaded = torch_load(CKPTDIR / "clip" / "open_clip_pytorch_model.bin")
loaded.update({"attn_mask": clip_encoder.attn_mask, "mean": clip_encoder.mean, "std": clip_encoder.std})
load_state_dict(clip_encoder, loaded)
Device.DEFAULT=original_device
@TinyJit
def denoise_step(x:Tensor, x_x:Tensor, t_t:Tensor, uc_c:Tensor, sqrt_alphas_cumprod_t:Tensor, sqrt_one_minus_alphas_cumprod_t:Tensor,
alpha_prev:Tensor, unet:UNetModel, GPUS) -> Tensor:
out_uncond, out = unet(x_x, t_t, uc_c).to("CPU").reshape(-1, 2, 4, 64, 64).chunk(2, dim=1)
out_uncond = out_uncond.squeeze(1).shard(GPUS,axis=0)
out = out.squeeze(1).shard(GPUS,axis=0)
v_t = out_uncond + 8.0 * (out - out_uncond)
e_t = sqrt_alphas_cumprod_t * v_t + sqrt_one_minus_alphas_cumprod_t * x
pred_x0 = sqrt_alphas_cumprod_t * x - sqrt_one_minus_alphas_cumprod_t * v_t
dir_xt = (1. - alpha_prev).sqrt() * e_t
x_prev = alpha_prev.sqrt() * pred_x0 + dir_xt
return x_prev.realize()
def shard_tensor(t:Tensor) -> Tensor: return t.shard(GPUS, axis=0) if len(GPUS) > 1 else t.to(GPUS[0])
def get_batch(whole:Tensor, i:int, bs:int) -> tuple[Tensor, int]:
batch = whole[i: i + bs].to("CPU")
if (unpadded_bs:=batch.shape[0]) < bs:
batch = batch.cat(batch[-1:].expand(bs - unpadded_bs, *batch[-1].shape))
return batch, unpadded_bs
@Tensor.train(mode=False)
def eval_unet(eval_inputs:list[dict], unet:UNetModel, cond_stage:FrozenOpenClipEmbedder, first_stage:AutoencoderKL,
inception:FidInceptionV3, clip:OpenClipEncoder) -> tuple[float, float]:
# Eval is divided into 5 jits, one per model
# It doesn't make sense to merge these jits, e.g. unet repeats 50 times in isolation; images fork to separate inception/clip
# We're generating and scoring 30,000 images per eval, and all the data can flow through one jit at a time
# To maximize throughput for each jit, we have only one model/jit on the GPU at a time, and pool outputs from each jit off-GPU
for model in (unet, first_stage, inception, clip):
Tensor.realize(*[p.to_("CPU") for p in get_parameters(model)])
uc_written = False
models = (cond_stage, unet, first_stage, inception, clip)
jits = (jit_context:=TinyJit(cond_stage.embed_tokens), denoise_step, vae_decode, jit_inception:=TinyJit(inception),
jit_clip:=TinyJit(clip.get_clip_score))
all_bs = (CONTEXT_BS, DENOISE_BS, DECODE_BS, INCEPTION_BS, CLIP_BS)
if (EVAL_SAMPLES:=getenv("EVAL_SAMPLES", 0)) and EVAL_SAMPLES > 0:
eval_inputs = eval_inputs[0:EVAL_SAMPLES]
output_shapes = [(ns:=len(eval_inputs),77), (ns,77,1024), (ns,4,64,64), (ns,3,512,512), (ns,2048), (ns,)]
# Writing progress to disk lets us resume eval if we crash
stages = ["tokens", "embeds", "latents", "imgs", "inception", "clip"]
disk_tensor_names, disk_tensor_shapes = stages + ["end", "uc"], output_shapes + [(6,), (1,77,1024)]
if not all(os.path.exists(f"{EVAL_CKPT_DIR}/{name}.bytes") for name in disk_tensor_names):
for name, shape in zip(disk_tensor_names, disk_tensor_shapes):
file = Path(f"{EVAL_CKPT_DIR}/{name}.bytes")
file.unlink(missing_ok=True)
with file.open("wb") as f: f.truncate(prod(shape) * 4)
progress = {name: Tensor.empty(*shape, device=f"disk:{EVAL_CKPT_DIR}/{name}.bytes", dtype=dtypes.int if name in {"tokens", "end"} else dtypes.float)
for name, shape in zip(disk_tensor_names, disk_tensor_shapes)}
def embed_tokens(tokens:Tensor) -> Tensor:
nonlocal uc_written
if not uc_written:
with Context(BEAM=0): progress["uc"].assign(cond_stage.embed_tokens(cond_stage.tokenize("").to(GPUS)).to("CPU").realize()).realize()
uc_written = True
return jit_context(shard_tensor(tokens))
def generate_latents(embeds:Tensor) -> Tensor:
uc_c = Tensor.stack(progress["uc"].to("CPU").expand(bs, 77, 1024), embeds, dim=1).reshape(-1, 77, 1024)
uc_c = shard_tensor(uc_c)
x = shard_tensor(Tensor.randn(bs,4,64,64))
for step_idx, timestep in enumerate(tqdm(eval_timesteps)):
reversed_idx = Tensor([50 - step_idx - 1], device=GPUS)
alpha_prev = eval_alphas_prev[reversed_idx]
ts = Tensor.full(bs, fill_value=timestep, dtype=dtypes.int, device="CPU")
ts_ts = shard_tensor(ts.cat(ts))
ts = shard_tensor(ts)
sqrt_alphas_cumprod_t = sqrt_alphas_cumprod[ts].reshape(bs, 1, 1, 1)
sqrt_one_minus_alphas_cumprod_t = sqrt_one_minus_alphas_cumprod[ts].reshape(bs, 1, 1, 1)
x_x = shard_tensor(Tensor.stack(x.to("CPU"), x.to("CPU"), dim=1).reshape(-1, 4, 64, 64))
x.assign(denoise_step(x, x_x, ts_ts, uc_c, sqrt_alphas_cumprod_t, sqrt_one_minus_alphas_cumprod_t, alpha_prev, unet, GPUS)).realize()
return x
def decode_latents(latents:Tensor) -> Tensor: return vae_decode(shard_tensor(latents), first_stage, disable_beam=True)
def generate_inception(imgs:Tensor) -> Tensor: return jit_inception(shard_tensor(imgs))[:,:,0,0]
def calc_clip_scores(batch:Tensor, batch_tokens:Tensor) -> Tensor:
# Tensor.interpolate does not yet support bicubic, so we use PIL
batch = (batch.to(GPUS[0]).permute(0,2,3,1) * 255).clip(0, 255).cast(dtypes.uint8).numpy()
batch = [np.array(PIL.Image.fromarray(batch[i]).resize((224,224), PIL.Image.BICUBIC)) for i in range(bs)]
batch = shard_tensor(Tensor(np.stack(batch, axis=0).transpose(0,3,1,2), device="CPU").realize())
batch = batch.cast(dtypes.float) / 255
batch = (batch - model.mean) / model.std
batch = jit_clip(shard_tensor(batch_tokens), batch)
return batch
callbacks = (embed_tokens, generate_latents, decode_latents, generate_inception, calc_clip_scores)
# save every forward pass output to disk; NOTE: this needs ~100 GB disk space because 30k images are large
def stage_progress(stage_idx:int) -> int: return progress["end"].to("CPU")[stage_idx].item()
if stage_progress(0) < len(eval_inputs):
tokens = []
for i in tqdm(range(0, len(eval_inputs), CONTEXT_BS)):
subset = [cond_stage.tokenize(row["caption"], device="CPU") for row in eval_inputs[i: i+CONTEXT_BS]]
tokens.append(Tensor.cat(*subset, dim=0).realize())
progress["tokens"].assign(Tensor.cat(*tokens, dim=0).realize()).realize()
progress["end"][0:1].assign(Tensor([len(eval_inputs)], dtype=dtypes.int)).realize()
prev_stage = "tokens"
tokens = progress["tokens"]
# wrapper code for every model
for stage_idx, model, jit, bs, callback in zip(range(1,6), models, jits, all_bs, callbacks):
stage = stages[stage_idx]
if stage_progress(stage_idx) >= len(eval_inputs):
prev_stage = stage
continue # use cache
t0 = time.perf_counter()
print(f"starting eval with model: {model}")
if stage_idx == 1: inputs = tokens
elif stage_idx == 5: inputs = progress["imgs"]
else: inputs = progress[prev_stage]
Tensor.realize(*[p.to_(GPUS) for p in get_parameters(model)])
for batch_idx in tqdm(range(stage_progress(stage_idx), inputs.shape[0], bs)):
t1 = time.perf_counter()
batch, unpadded_bs = get_batch(inputs, batch_idx, bs)
if isinstance(model, OpenClipEncoder): batch = callback(batch, get_batch(tokens, batch_idx, bs)[0].realize())
else: batch = callback(batch)
# to(GPUS[0]) is necessary for this to work, without that the result is still on GPUS, probably due to a bug
batch = batch.to(GPUS[0]).to("CPU")[0:unpadded_bs].realize()
progress[stage][batch_idx: batch_idx + bs].assign(batch).realize()
# keep track of what our last output was, so we can resume from there if we crash in this loop
progress["end"][stage_idx: stage_idx + 1].assign(Tensor([batch_idx + bs], dtype=dtypes.int)).realize()
print(f"model: {model}, batch_idx: {batch_idx}, elapsed: {(time.perf_counter() - t1):.2f}")
del batch
jit.reset()
Tensor.realize(*[p.to_("CPU") for p in get_parameters(model)])
print(f"done with model: {model}, elapsed: {(time.perf_counter() - t0):.2f}")
prev_stage = stage
inception_stats_fn = str(DATADIR / "coco2014" / "val2014_30k_stats.npz")
fid_score = inception.compute_score(progress["inception"].to("CPU"), inception_stats_fn)
clip_score = progress["clip"].to(GPUS[0]).mean().item()
for name in disk_tensor_names:
Path(f"{EVAL_CKPT_DIR}/{name}.bytes").unlink(missing_ok=True)
if EVAL_SAMPLES and BEAM:
print("BEAM COMPLETE", flush=True) # allows wrapper script to detect BEAM search completion and retry if it failed
sys.exit() # Don't eval additional models; we don't care about clip/fid scores when running BEAM on eval sample subset
return clip_score, fid_score
# evaluate checkpoints in reverse chronological order
for ckpt_iteration, p in sorted(eval_queue, reverse=True):
unet_ckpt = safe_load(p)
load_state_dict(unet, unet_ckpt)
clip_score, fid_score = eval_unet(eval_inputs, unet, model.cond_stage_model, model.first_stage_model, inception, clip_encoder)
converged = True if clip_score >= 0.15 and fid_score <= 90 else False
print(f"eval results for {EVAL_CKPT_DIR}/{p.name}: clip={clip_score}, fid={fid_score}, converged={converged}")
if WANDB:
wandb.log({"eval/ckpt_iteration": ckpt_iteration, "eval/clip_score": clip_score, "eval/fid_score": fid_score})
if converged and STOP_IF_CONVERGED:
print(f"Convergence detected, exiting early before evaluating other checkpoints due to STOP_IF_CONVERGED={STOP_IF_CONVERGED}")
sys.exit()
# for testing
return clip_score, fid_score, ckpt_iteration
if __name__ == "__main__":
# inference only
Tensor.training = False
+5 -142
View File
@@ -3,7 +3,7 @@ from pathlib import Path
import multiprocessing
from tinygrad import Device, GlobalCounters, Tensor, TinyJit, dtypes
from tinygrad.helpers import getenv, BEAM, WINO, round_up, diskcache_clear, Profiling
from tinygrad.helpers import getenv, BEAM, WINO, round_up, diskcache_clear, FUSE_CONV_BW, Profiling
from tinygrad.nn.state import get_parameters, get_state_dict, load_state_dict, safe_load, safe_save
from tinygrad.nn.optim import LAMB, LARS, SGD, OptimizerGroup, Adam, AdamW
@@ -707,7 +707,7 @@ def train_unet3d():
```BASEDIR=<folder_path> ./examples/mlperf/scripts/setup_kits19_dataset.sh```
2) To start training the model, run the following:
```time PYTHONPATH=. WANDB=1 TRAIN_BEAM=3 GPUS=6 BS=6 MODEL=unet3d python3 examples/mlperf/model_train.py```
```time PYTHONPATH=. WANDB=1 TRAIN_BEAM=3 FUSE_CONV_BW=1 GPUS=6 BS=6 MODEL=unet3d python3 examples/mlperf/model_train.py```
"""
from examples.mlperf.losses import dice_ce_loss
from examples.mlperf.metrics import dice_score
@@ -749,6 +749,7 @@ def train_unet3d():
"train_beam": TRAIN_BEAM,
"eval_beam": EVAL_BEAM,
"wino": WINO.value,
"fuse_conv_bw": FUSE_CONV_BW.value,
"gpus": GPUS,
"default_float": dtypes.default_float.name
}
@@ -1308,7 +1309,7 @@ def train_llama3():
EVAL_BS = config["EVAL_BS"] = getenv("EVAL_BS", 16)
EVAL_TARGET = config["EVAL_TARGET"] = getenv("EVAL_TARGET", 5.6)
# LR=1e-4 TRAIN_ON_VAL=1 DEFAULT_FLOAT=bfloat16 JITBEAM=2 OPTIM_DTYPE=bfloat16 LLAMA3_SIZE=1B WARMUP_STEPS=36 DECAY_STEPS=360 SEQLEN=512 PYTHONPATH=. AMD=1 AMD_LLVM=0 MODEL=llama3 python3 examples/mlperf/model_train.py
# LR=1e-4 TRAIN_ON_VAL=1 DEFAULT_FLOAT=bfloat16 FUSE_ARANGE=1 JITBEAM=2 OPTIM_DTYPE=bfloat16 LLAMA3_SIZE=1B WARMUP_STEPS=36 DECAY_STEPS=360 SEQLEN=512 PYTHONPATH=. AMD=1 AMD_LLVM=0 MODEL=llama3 python3 examples/mlperf/model_train.py
# trains to 7
opt_adamw_beta_1 = 0.9
@@ -1492,144 +1493,6 @@ def train_llama3():
safe_save(get_state_dict(model), fn)
break
def train_stable_diffusion():
from extra.models.unet import UNetModel
from examples.mlperf.dataloader import batch_load_train_stable_diffusion
from examples.mlperf.lr_schedulers import LambdaLR, LambdaLinearScheduler
from examples.mlperf.initializers import init_stable_diffusion
from examples.mlperf.helpers import get_training_state
import numpy as np
config = {}
GPUS = config["GPUS"] = [f"{Device.DEFAULT}:{i}" for i in range(getenv("GPUS", 1))]
seed = config["seed"] = getenv("SEED", 12345)
# ** hyperparameters **
BS = config["BS"] = getenv("BS", 1 * len(GPUS))
BASE_LR = config["LEARNING_RATE"] = getenv("LEARNING_RATE", 2.5e-7)
# https://github.com/mlcommons/training_policies/blob/cfa99da479b8d5931f7a3c67612d021dfb47510a/training_rules.adoc#benchmark_specific_rules
# "Checkpoint must be collected every 512,000 images. CEIL(512000 / global_batch_size) if 512000 is not divisible by GBS."
# NOTE: It's inferred that "steps" is the unit for the output of the CEIL formula, based on all other cases of CEIL in the rules
CKPT_STEP_INTERVAL = config["CKPT_STEP_INTERVAL"] = getenv("CKPT_STEP_INTERVAL", math.ceil(512_000 / BS))
CKPTDIR = config["CKPTDIR"] = Path(getenv("CKPTDIR", "./checkpoints"))
DATADIR = config["DATADIR"] = Path(getenv("DATADIR", "./datasets"))
UNET_CKPTDIR = config["UNET_CKPTDIR"] = Path(getenv("UNET_CKPTDIR", "./checkpoints"))
TOTAL_CKPTS = config["TOTAL_CKPTS"] = getenv("TOTAL_CKPTS", 0)
print(f"training on {GPUS}")
lr = BS * BASE_LR
print(f"BS={BS}, BASE_LR={BASE_LR}, lr={lr}")
print(f"CKPT_STEP_INTERVAL = {CKPT_STEP_INTERVAL}")
for x in GPUS: Device[x]
if (WANDB := getenv("WANDB", "")):
import wandb
wandb.init(config=config, project="MLPerf-Stable-Diffusion")
Tensor.manual_seed(seed) # seed for weight initialization
model, unet, sqrt_alphas_cumprod, sqrt_one_minus_alphas_cumprod = init_stable_diffusion("v2-mlperf-train", CKPTDIR / "sd" / "512-base-ema.ckpt", GPUS)
optimizer = AdamW(get_parameters(unet))
lambda_lr_callback = LambdaLinearScheduler(1000, 1.0, 1.0, 1e-06, 10000000000000).schedule
lr_scheduler = LambdaLR(optimizer, Tensor(lr, dtype=dtypes.float, device=optimizer.device), lambda_lr_callback)
@TinyJit
def train_step(mean:Tensor, logvar:Tensor, tokens:Tensor, unet:UNetModel, optimizer:LAMB, lr_scheduler:LambdaLR) -> Tensor:
optimizer.zero_grad()
timestep = Tensor.randint(BS, low=0, high=model.alphas_cumprod.shape[0], dtype=dtypes.int, device=GPUS[0])
latent_randn = Tensor.randn(*mean.shape, device=GPUS[0])
noise = Tensor.randn(*mean.shape, device=GPUS[0])
for t in (mean, logvar, tokens, timestep, latent_randn, noise):
t.shard_(GPUS, axis=0)
std = Tensor.exp(0.5 * logvar.clamp(-30.0, 20.0))
latent = (mean + std * latent_randn) * 0.18215
sqrt_alphas_cumprod_t = sqrt_alphas_cumprod[timestep].reshape(timestep.shape[0], 1, 1, 1)
sqrt_one_minus_alphas_cumprod_t = sqrt_one_minus_alphas_cumprod[timestep].reshape(timestep.shape[0], 1, 1, 1)
latent_with_noise = sqrt_alphas_cumprod_t * latent + sqrt_one_minus_alphas_cumprod_t * noise
v_true = sqrt_alphas_cumprod_t * noise - sqrt_one_minus_alphas_cumprod_t * latent
context = model.cond_stage_model.embed_tokens(tokens)
out = unet(latent_with_noise, timestep, context)
loss = ((out - v_true) ** 2).mean()
del mean, logvar, std, latent, noise, sqrt_alphas_cumprod_t, sqrt_one_minus_alphas_cumprod_t
del out, v_true, context, latent_randn, tokens, timestep
loss.backward()
optimizer.step()
lr_scheduler.step()
loss, out_lr = loss.detach().to("CPU"), optimizer.lr.to("CPU")
Tensor.realize(loss, out_lr)
return loss, out_lr
# checkpointing takes ~9 minutes without this, and ~1 minute with this
@TinyJit
def ckpt_to_cpu():
ckpt = get_training_state(unet, optimizer, lr_scheduler)
# move to CPU first so more GPU bufs aren't created (can trigger OOM)
for k,v in ckpt.items(): ckpt[k] = v.detach().to("CPU")
Tensor.realize(*[v for v in ckpt.values()])
for k,v in ckpt.items(): ckpt[k] = v.cast(v.dtype.base).contiguous()
Tensor.realize(*[v for v in ckpt.values()])
return ckpt
# training loop
dl = batch_load_train_stable_diffusion(f'{DATADIR}/laion-400m/webdataset-moments-filtered/{{00000..00831}}.tar', BS)
# for tests
saved_checkpoints = []
train_start_time = time.perf_counter()
t0 = t6 = time.perf_counter()
for i, batch in enumerate(dl, start=1):
loop_time = time.perf_counter() - t0
t0 = time.perf_counter()
dl_time = t0 - t6
GlobalCounters.reset()
mean, logvar = np.split(np.concatenate(batch["npy"], axis=0), 2, axis=1)
mean, logvar = Tensor(mean, dtype=dtypes.float32, device="CPU"), Tensor(logvar, dtype=dtypes.float32, device="CPU")
tokens = []
for text in batch['txt']: tokens += model.cond_stage_model.tokenizer.encode(text, pad_with_zeros=True)
tokens = Tensor(tokens, dtype=dtypes.int32, device="CPU").reshape(-1, 77)
t1 = time.perf_counter()
loss, lr = train_step(mean, logvar, tokens, unet, optimizer, lr_scheduler)
loss_item, lr_item = loss.item(), lr.item()
t2 = time.perf_counter()
if i == 3:
for _ in range(3): ckpt_to_cpu() # do this at the beginning of run to prevent OOM surprises when checkpointing
print("BEAM COMPLETE", flush=True) # allows wrapper script to detect BEAM search completion and retry if it failed
total_train_time = time.perf_counter() - train_start_time
if WANDB:
wandb.log({"train/loss": loss_item, "train/lr": lr_item, "train/loop_time_prev": loop_time, "train/dl_time": dl_time, "train/step": i,
"train/GFLOPS": GlobalCounters.global_ops * 1e-9 / (t2-t1), "train/input_prep_time": t1-t0,
"train/train_step_time": t2-t1, "train/total_time": total_train_time})
if i == 1 and wandb.run is not None:
with open(f"{UNET_CKPTDIR}/wandb_run_id_{wandb.run.id}", "w") as f:
f.write(f"wandb.run.id = {wandb.run.id}")
if i % CKPT_STEP_INTERVAL == 0:
# https://github.com/mlcommons/training_policies/blob/cfa99da479b8d5931f7a3c67612d021dfb47510a/training_rules.adoc#benchmark_specific_rules
# "evaluation is done offline, the time is not counted towards the submission time."
fn = f"{UNET_CKPTDIR}/{i}.safetensors"
print(f"saving unet checkpoint at {fn}")
saved_checkpoints.append(fn)
safe_save({k.replace("model.", ""):v for k,v in ckpt_to_cpu().items() if k.startswith("model.")}, fn)
if TOTAL_CKPTS and i == TOTAL_CKPTS * CKPT_STEP_INTERVAL:
print(f"ending run after {i} steps ({TOTAL_CKPTS} checkpoints collected)")
return saved_checkpoints
t3 = time.perf_counter()
print(f"""step {i}: {GlobalCounters.global_ops * 1e-9 / (t2-t1):9.2f} GFLOPS, mem_used: {GlobalCounters.mem_used / 1e9:.2f} GB,
loop_time_prev: {loop_time:.2f}, dl_time: {dl_time:.2f}, input_prep_time: {t1-t0:.2f}, train_step_time: {t2-t1:.2f},
t3-t2: {t3-t2:.4f}, loss:{loss_item:.5f}, lr:{lr_item:.3e}, total_train_time:{total_train_time:.2f}
""")
t6 = time.perf_counter()
if __name__ == "__main__":
multiprocessing.set_start_method('spawn')
@@ -1638,7 +1501,7 @@ if __name__ == "__main__":
else: bench_log_manager = contextlib.nullcontext()
with Tensor.train():
for m in getenv("MODEL", "resnet,retinanet,unet3d,rnnt,bert,maskrcnn,stable_diffusion").split(","):
for m in getenv("MODEL", "resnet,retinanet,unet3d,rnnt,bert,maskrcnn").split(","):
nm = f"train_{m}"
if nm in globals():
print(f"training {m}")
@@ -1,72 +0,0 @@
#!/usr/bin/env bash
DATETIME=${2:-$(date "+%m%d%H%M")}
LOGFILE="${HOME}/logs/sd_mi300x_${DATETIME}.log"
# UNET_CKPTDIR must be set: training saves checkpoints to this path, then a separate eval process scans this path to know which checkpoints to eval
export UNET_CKPTDIR="${HOME}/stable_diffusion/training_checkpoints/${DATETIME}"
mkdir -p "${HOME}/logs" "$UNET_CKPTDIR"
# run this script in isolation when using the --bg flag
if [[ "${1:-}" == "--bg" ]]; then
echo "logging output to $LOGFILE"
echo "saving UNet checkpoints to $UNET_CKPTDIR"
script_path="$(readlink -f "${BASH_SOURCE[0]}")"
nohup bash "$script_path" run "$DATETIME" >"$LOGFILE" 2>&1 & disown $!
exit 0
fi
# venv management
if [[ -d .venv-sd-mlperf ]]; then
. .venv-sd-mlperf/bin/activate
else
python3 -m venv .venv-sd-mlperf && . .venv-sd-mlperf/bin/activate
pip install --index-url https://download.pytorch.org/whl/cpu torch && pip install tqdm numpy ftfy regex pillow scipy wandb webdataset
fi
pip list
apt list --installed | grep amdgpu
rocm-smi --version
modinfo amdgpu | grep version
export BEAM=2 BEAM_UOPS_MAX=8000 BEAM_UPCAST_MAX=256 BEAM_LOCAL_MAX=1024 BEAM_MIN_PROGRESS=5 IGNORE_JIT_FIRST_BEAM=1 HCQDEV_WAIT_TIMEOUT_MS=300000
export AMD_LLVM=0 # bf16 seems to require this
export DATADIR="/raid/datasets/stable_diffusion"
export CKPTDIR="/raid/weights/stable_diffusion"
export EVAL_CKPT_DIR=$UNET_CKPTDIR
export MODEL="stable_diffusion" PYTHONPATH="."
export GPUS=8 BS=304
export CONTEXT_BS=816 DENOISE_BS=600 DECODE_BS=384 INCEPTION_BS=560 CLIP_BS=240
export WANDB=1
export PARALLEL=4
export PYTHONUNBUFFERED=1
sudo rocm-smi -d 0 1 2 3 4 5 6 7 --setperfdeterminism 1500 || exit 1
# Retry BEAM search if script fails before BEAM COMPLETE is printed, but don't retry after that
run_retry(){ local try=0 max=5 code tmp py pgid kids
while :; do
tmp=$(mktemp)
setsid bash -c 'exec env "$@"' _ "$@" > >(tee -a "$LOGFILE" | tee "$tmp") 2>&1 &
py=$!; pgid=$(ps -o pgid= -p "$py" | tr -d ' ')
wait "$py"; code=$?
[[ -n "$pgid" ]] && { kill -TERM -"$pgid" 2>/dev/null; sleep 1; kill -KILL -"$pgid" 2>/dev/null; }
kids=$(pgrep -P "$py" || true)
while [[ -n "$kids" ]]; do
kill -TERM $kids 2>/dev/null; sleep 0.5
kids=$(for k in $kids; do pgrep -P "$k" || true; done)
done
grep -q 'BEAM COMPLETE' "$tmp" && { rm -f "$tmp"; return 1; }
rm -f "$tmp"
((code==0)) && return 0
((try>=max)) && return 2
((try++)); sleep 90; echo "try = ${try}"
done
}
# Power limiting to 400W is only needed if GPUs fall out of sync (causing 2.2x increased train time) at higher power, which has been observed at 450W
sudo rocm-smi -d 0 1 2 3 4 5 6 7 --setpoweroverdrive 750 && \
run_retry TOTAL_CKPTS=7 python3 examples/mlperf/model_train.py; (( $? == 2 )) && { echo "training failed before BEAM completion"; exit 2; }
sleep 90
run_retry EVAL_SAMPLES=600 python3 examples/mlperf/model_eval.py; (( $? == 2 )) && { echo "eval failed before BEAM completion"; exit 2; }
# Checkpoints will be evaluated in reverse chronological order, even if above training crashed early
# STOP_IF_CONVERGED=1: Stop the eval after the first time convergence is detected; no more checkpoints will be evaluated after that.
STOP_IF_CONVERGED=1 python3 examples/mlperf/model_eval.py
+3 -12
View File
@@ -1,4 +1,4 @@
import os, sys, pickle, time, re
import os, sys, pickle, time
import numpy as np
if "FLOAT16" not in os.environ: os.environ["FLOAT16"] = "1"
if "IMAGE" not in os.environ: os.environ["IMAGE"] = "2"
@@ -10,7 +10,7 @@ from tinygrad.helpers import DEBUG, getenv
from tinygrad.engine.realize import CompiledRunner
import onnx
from tinygrad.nn.onnx import OnnxRunner
from tinygrad.frontend.onnx import OnnxRunner
OPENPILOT_MODEL = sys.argv[1] if len(sys.argv) > 1 else "https://github.com/commaai/openpilot/raw/v0.9.7/selfdrive/modeld/models/supercombo.onnx"
OUTPUT = sys.argv[2] if len(sys.argv) > 2 else "/tmp/openpilot.pkl"
@@ -52,8 +52,6 @@ def compile(onnx_file):
kernel_count += 1
read_image_count += ei.prg.p.src.count("read_image")
gated_read_image_count += ei.prg.p.src.count("?read_image")
for v in [m.group(1) for m in re.finditer(r'(val\d+)\s*=\s*read_imagef\(', ei.prg.p.src)]:
if len(re.findall(fr'[\?\:]{v}\.[xyzw]', ei.prg.p.src)) > 0: gated_read_image_count += 1
print(f"{kernel_count=}, {read_image_count=}, {gated_read_image_count=}")
if (allowed_kernel_count:=getenv("ALLOWED_KERNEL_COUNT", -1)) != -1:
assert kernel_count == allowed_kernel_count, f"different kernels! {kernel_count=}, {allowed_kernel_count=}"
@@ -79,20 +77,13 @@ def test_vs_compile(run, new_inputs, test_val=None):
**{k:Tensor(v, device="NPY").realize() for k,v in new_inputs_numpy.items() if 'img' not in k}}
# run 20 times
step_times = []
for _ in range(20):
st = time.perf_counter()
out = run(**inputs)
mt = time.perf_counter()
val = out.numpy()
et = time.perf_counter()
step_times.append((et-st)*1e3)
print(f"enqueue {(mt-st)*1e3:6.2f} ms -- total run {step_times[-1]:6.2f} ms")
if (assert_time:=getenv("ASSERT_MIN_STEP_TIME")):
min_time = min(step_times)
assert min_time < assert_time, f"Speed regression, expected min step time of < {assert_time} ms but took: {min_time} ms"
print(f"enqueue {(mt-st)*1e3:6.2f} ms -- total run {(et-st)*1e3:6.2f} ms")
print(out, val.shape, val.dtype)
if test_val is not None: np.testing.assert_equal(test_val, val)
print("**** test done ****")
+3 -3
View File
@@ -1,8 +1,8 @@
import sys
from tinygrad import Tensor, fetch, GlobalCounters, dtypes
from tinygrad.uop.ops import UOp
from tinygrad.nn.onnx import OnnxRunner
from tinygrad.schedule.rangeify import get_rangeify_map
from tinygrad.frontend.onnx import OnnxRunner
from tinygrad.schedule.kernelize import get_kernelize_map
from tinygrad.engine.schedule import create_schedule_with_vars
from tinygrad.engine.realize import run_schedule
@@ -33,7 +33,7 @@ if __name__ == "__main__":
if not in_target_path[s]:
independent_set[s] = None
independent = UOp.sink(*independent_set.keys())
kernelized = get_rangeify_map(independent)
kernelized = get_kernelize_map(independent)
independent = independent.substitute(kernelized)
schedule, var_vals = create_schedule_with_vars(independent)
run_schedule(schedule)
@@ -27,7 +27,7 @@ class Model(nn.Module):
if __name__ == "__main__":
if getenv("TINY_BACKEND"):
import tinygrad.nn.torch # noqa: F401
import tinygrad.frontend.torch # noqa: F401
device = torch.device("tiny")
else:
device = torch.device({"METAL":"mps","NV":"cuda"}.get(Device.DEFAULT, "cpu"))
+8 -46
View File
@@ -9,13 +9,11 @@ from typing import Dict, Any
from PIL import Image
import numpy as np
from tinygrad import Device, GlobalCounters, dtypes, Tensor, TinyJit
from tinygrad.helpers import Timing, Context, getenv, fetch, colored, tqdm, flatten
from tinygrad.helpers import Timing, Context, getenv, fetch, colored, tqdm
from tinygrad.nn import Conv2d, GroupNorm
from tinygrad.nn.state import torch_load, load_state_dict, get_state_dict
from extra.models.clip import Closed, Tokenizer, FrozenOpenClipEmbedder
from extra.models import unet, clip
from extra.models.clip import Closed, Tokenizer
from extra.models.unet import UNetModel
from examples.mlperf.initializers import AutocastLinear, AutocastConv2d, AutocastGroupNorm, AutocastLayerNorm, zero_module, attn_f32_softmax, gelu_erf
from extra.bench_log import BenchEvent, WallTimeEvent
class AttnBlock:
@@ -156,46 +154,12 @@ unet_params: Dict[str,Any] = {
"use_linear": False,
}
mlperf_params: Dict[str,Any] = {"adm_in_ch": None, "in_ch": 4, "out_ch": 4, "model_ch": 320, "attention_resolutions": [4, 2, 1], "num_res_blocks": 2,
"channel_mult": [1, 2, 4, 4], "d_head": 64, "transformer_depth": [1, 1, 1, 1], "ctx_dim": 1024, "use_linear": True,
"num_groups":16, "st_norm_eps":1e-6}
class StableDiffusion:
def __init__(self, version:str|None=None, pretrained:str|None=None):
def __init__(self):
self.alphas_cumprod = get_alphas_cumprod()
if version != "v2-mlperf-train":
self.first_stage_model = AutoencoderKL() # only needed for decoding generated latents to images; not needed in mlperf training from preprocessed moments
if not version:
self.cond_stage_model = namedtuple("CondStageModel", ["transformer"])(transformer = namedtuple("Transformer", ["text_model"])(text_model = Closed.ClipTextTransformer()))
unet_init_params = unet_params
elif version in {"v2-mlperf-train", "v2-mlperf-eval"}:
unet_init_params = mlperf_params
clip.gelu = gelu_erf
self.cond_stage_model = FrozenOpenClipEmbedder(**{"dims": 1024, "n_heads": 16, "layers": 24, "return_pooled": False, "ln_penultimate": True,
"clip_tokenizer_version": "sd_mlperf_v5_0"})
unet.Linear, unet.Conv2d, unet.GroupNorm, unet.LayerNorm = AutocastLinear, AutocastConv2d, AutocastGroupNorm, AutocastLayerNorm
unet.attention, unet.gelu, unet.mixed_precision_dtype = attn_f32_softmax, gelu_erf, dtypes.bfloat16
if pretrained:
print("loading text encoder")
weights: dict[str,Tensor] = {k.replace("cond_stage_model.", "", 1):v for k,v in torch_load(pretrained)["state_dict"].items() if k.startswith("cond_stage_model.")}
weights["model.attn_mask"] = Tensor.full((77, 77), fill_value=float("-inf")).triu(1)
load_state_dict(self.cond_stage_model, weights)
# only the eval model needs the decoder
if version == "v2-mlperf-eval":
print("loading image latent encoder")
weights = {k.replace("first_stage_model.", "", 1):v for k,v in torch_load(pretrained)["state_dict"].items() if k.startswith("first_stage_model.")}
load_state_dict(self.first_stage_model, weights)
self.model = namedtuple("DiffusionModel", ["diffusion_model"])(diffusion_model = UNetModel(**unet_init_params))
if version == "v2-mlperf-train":
# the mlperf reference inits certain weights as zeroes
for bb in flatten(self.model.diffusion_model.input_blocks) + self.model.diffusion_model.middle_block + flatten(self.model.diffusion_model.output_blocks):
if isinstance(bb, unet.ResBlock):
zero_module(bb.out_layers[3])
elif isinstance(bb, unet.SpatialTransformer):
zero_module(bb.proj_out)
zero_module(self.model.diffusion_model.out[2])
self.model = namedtuple("DiffusionModel", ["diffusion_model"])(diffusion_model = UNetModel(**unet_params))
self.first_stage_model = AutoencoderKL()
self.cond_stage_model = namedtuple("CondStageModel", ["transformer"])(transformer = namedtuple("Transformer", ["text_model"])(text_model = Closed.ClipTextTransformer()))
def get_x_prev_and_pred_x0(self, x, e_t, a_t, a_prev):
temperature = 1
@@ -269,14 +233,12 @@ if __name__ == "__main__":
# load in weights
with WallTimeEvent(BenchEvent.LOAD_WEIGHTS):
load_state_dict(model, torch_load(fetch('https://huggingface.co/CompVis/stable-diffusion-v-1-4-original/resolve/main/sd-v1-4.ckpt', 'sd-v1-4.ckpt'))['state_dict'], verbose=False, strict=False, realize=False)
load_state_dict(model, torch_load(fetch('https://huggingface.co/CompVis/stable-diffusion-v-1-4-original/resolve/main/sd-v1-4.ckpt', 'sd-v1-4.ckpt'))['state_dict'], strict=False)
if args.fp16:
for k,v in get_state_dict(model).items():
if k.startswith("model"):
v.replace(v.cast(dtypes.float16))
Tensor.realize(*get_state_dict(model).values())
v.replace(v.cast(dtypes.float16).realize())
# run through CLIP to get context
tokenizer = Tokenizer.ClipTokenizer()
+1 -1
View File
@@ -32,7 +32,7 @@ if __name__ == "__main__":
lr = 5e-3
transform = ComposeTransforms([
lambda x: [Image.fromarray(xx).resize((64, 64)) for xx in x],
lambda x: [Image.fromarray(xx, mode='L').resize((64, 64)) for xx in x],
lambda x: np.stack([np.asarray(xx) for xx in x], 0),
lambda x: x / 255.0,
lambda x: np.tile(np.expand_dims(x, 1), (1, 3, 1, 1)).astype(np.float32),
+1 -1
View File
@@ -109,7 +109,7 @@ class TextDecoder:
def forward(self, x:Tensor, pos:Union[Variable, Literal[0]], encoded_audio:Tensor):
seqlen = x.shape[-1]
x = self.token_embedding(x) + self.positional_embedding.shrink(((pos, pos+seqlen), None))
x = self.token_embedding(x) + self.positional_embedding.shrink(((pos, pos+seqlen), None, None))
for block in self.blocks: x = block(x, xa=encoded_audio, mask=self.mask, len=pos)
return self.output_tok(x)
+1 -1
View File
@@ -2,7 +2,7 @@
import os
from ultralytics import YOLO
from pathlib import Path
from tinygrad.nn.onnx import OnnxRunner
from tinygrad.frontend.onnx import OnnxRunner
from extra.onnx_helpers import get_example_inputs
os.chdir("/tmp")
+3 -2
View File
@@ -49,7 +49,8 @@ def rangeify_kernel3():
b = Tensor.empty(N,N)
c = a@b
#c = c.reshape((32,2,16,4,32,2,16,4)).contiguous()
sink = c.schedule()[-1].ast
with Context(RANGEIFY=1):
sink = c.schedule()[-1].ast
#print(sink)
opts = [Opt(OptOps.UPCAST, 0, 4), Opt(OptOps.LOCAL, 0, 16), Opt(OptOps.UPCAST, 0, 2)]
@@ -328,7 +329,7 @@ if __name__ == "__main__":
elif HL == 1: hprg = hl_spec_kernel3()
else: hprg = hand_spec_kernel3()
if HL == 3:
with Context(BLOCK_REORDER=0):
with Context(RANGEIFY=1, BLOCK_REORDER=0):
prg = get_program(hprg, Device.default.renderer)
else:
prg = get_program(hprg, Device.default.renderer)
+1
View File
@@ -7,6 +7,7 @@ bert_train_params = {
"GPUS": 6,
"BS": 96,
"EVAL_BS": 96,
"FUSE_ARANGE": 1,
"BASEDIR": "/raid/datasets/wiki",
}
+1 -1
View File
@@ -50,7 +50,7 @@ def ioctls_from_header():
hdr = (pathlib.Path(__file__).parent / "kfd_ioctl.h").read_text().replace("\\\n", "")
pattern = r'#define\s+(AMDKFD_IOC_[A-Z0-9_]+)\s+AMDKFD_IOW?R?\((0x[0-9a-fA-F]+),\s+struct\s([A-Za-z0-9_]+)\)'
matches = re.findall(pattern, hdr, re.MULTILINE)
return {int(nr, 0x10):(name, getattr(kfd_ioctl, "struct_"+sname, None)) for name, nr, sname in matches}
return {int(nr, 0x10):(name, getattr(kfd_ioctl, "struct_"+sname)) for name, nr, sname in matches}
nrs = ioctls_from_header()
@ctypes.CFUNCTYPE(ctypes.c_int, ctypes.c_int, ctypes.c_ulong, ctypes.c_void_p)
File diff suppressed because it is too large Load Diff
+1 -1
View File
@@ -1,7 +1,7 @@
import onnx, yaml, tempfile, time, argparse, json
from pathlib import Path
from typing import Any
from tinygrad.nn.onnx import OnnxRunner
from tinygrad.frontend.onnx import OnnxRunner
from extra.onnx_helpers import validate, get_example_inputs
from extra.huggingface_onnx.huggingface_manager import DOWNLOADS_DIR, snapshot_download_with_retry
+15 -32
View File
@@ -9,9 +9,6 @@ from PIL import Image
import numpy as np
import re, gzip
# Allow for monkeypatching for mlperf.
gelu = Tensor.gelu
@lru_cache()
def default_bpe():
# Clip tokenizer, taken from https://github.com/openai/CLIP/blob/main/clip/simple_tokenizer.py (MIT license)
@@ -56,8 +53,8 @@ class Tokenizer:
cs = [chr(n) for n in cs]
return dict(zip(bs, cs))
class ClipTokenizer:
def __init__(self, version=None):
self.byte_encoder, self.version = Tokenizer.bytes_to_unicode(), version
def __init__(self):
self.byte_encoder = Tokenizer.bytes_to_unicode()
merges = gzip.open(default_bpe()).read().decode("utf-8").split('\n')
merges = merges[1:49152-256-2+1]
merges = [tuple(merge.split()) for merge in merges]
@@ -65,17 +62,11 @@ class Tokenizer:
vocab = vocab + [v+'</w>' for v in vocab]
for merge in merges:
vocab.append(''.join(merge))
if self.version == "sd_mlperf_v5_0":
import regex
vocab.extend(['<start_of_text>', '<end_of_text>'])
self.cache = {'<start_of_text>': '<start_of_text>', '<end_of_text>': '<end_of_text>'}
self.pat = regex.compile(r"""<start_of_text>|<end_of_text>|'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+""", regex.IGNORECASE)
else:
vocab.extend(['<|startoftext|>', '<|endoftext|>'])
self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'}
self.pat = re.compile(r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[^\s]+""", re.IGNORECASE)
vocab.extend(['<|startoftext|>', '<|endoftext|>'])
self.encoder = dict(zip(vocab, range(len(vocab))))
self.bpe_ranks = dict(zip(merges, range(len(merges))))
self.cache = {'<|startoftext|>': '<|startoftext|>', '<|endoftext|>': '<|endoftext|>'}
self.pat = re.compile(r"""<\|startoftext\|>|<\|endoftext\|>|'s|'t|'re|'ve|'m|'ll|'d|[^\s]+""", re.IGNORECASE)
def bpe(self, token):
if token in self.cache:
@@ -119,17 +110,8 @@ class Tokenizer:
def encode(self, text:str, pad_with_zeros:bool=False) -> List[int]:
bpe_tokens: List[int] = []
if self.version == "sd_mlperf_v5_0":
import regex, ftfy, html
text = ftfy.fix_text(text)
text = html.unescape(html.unescape(text)).strip()
text = Tokenizer.whitespace_clean(text).lower()
re_module = regex
else:
text = Tokenizer.whitespace_clean(text.strip()).lower()
re_module = re
for token in re_module.findall(self.pat, text):
text = Tokenizer.whitespace_clean(text.strip()).lower()
for token in re.findall(self.pat, text):
token = ''.join(self.byte_encoder[b] for b in token.encode('utf-8'))
bpe_tokens.extend(self.encoder[bpe_token] for bpe_token in self.bpe(token).split(' '))
# Truncation, keeping two slots for start and end tokens.
@@ -270,8 +252,10 @@ class Open:
q,k,v = [y.reshape(T, B*self.n_heads, self.d_head).transpose(0, 1).reshape(B, self.n_heads, T, self.d_head) for y in proj.chunk(3)]
attn_output = Tensor.scaled_dot_product_attention(q, k, v, attn_mask=attn_mask)
attn_output = attn_output.permute(2, 0, 1, 3).reshape(T, B, C)
attn_output = attn_output.permute(2, 0, 1, 3).reshape(T*B, C)
attn_output = self.out_proj(attn_output)
attn_output = attn_output.reshape(T, B, C)
return attn_output
@@ -279,10 +263,9 @@ class Open:
def __init__(self, dims, hidden_dims):
self.c_fc = Linear(dims, hidden_dims)
self.c_proj = Linear(hidden_dims, dims)
self.gelu = gelu
def __call__(self, x:Tensor) -> Tensor:
return x.sequential([self.c_fc, self.gelu, self.c_proj])
return x.sequential([self.c_fc, Tensor.gelu, self.c_proj])
# https://github.com/mlfoundations/open_clip/blob/58e4e39aaabc6040839b0d2a7e8bf20979e4558a/src/open_clip/transformer.py#L210
class ResidualAttentionBlock:
@@ -367,15 +350,15 @@ class Open:
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/encoders/modules.py#L396
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/encoders/modules.py#L498
class FrozenOpenClipEmbedder(Embedder):
def __init__(self, dims:int, n_heads:int, layers:int, return_pooled:bool, ln_penultimate:bool=False, clip_tokenizer_version=None):
self.tokenizer = Tokenizer.ClipTokenizer(version=clip_tokenizer_version)
def __init__(self, dims:int, n_heads:int, layers:int, return_pooled:bool, ln_penultimate:bool=False):
self.tokenizer = Tokenizer.ClipTokenizer()
self.model = Open.ClipTextTransformer(dims, n_heads, layers)
self.return_pooled = return_pooled
self.input_key = "txt"
self.ln_penultimate = ln_penultimate
def tokenize(self, text:str, device:Optional[str]=None) -> Tensor:
return Tensor(self.tokenizer.encode(text, pad_with_zeros=True), dtype=dtypes.int32, device=device).reshape(1,-1)
return Tensor(self.tokenizer.encode(text, pad_with_zeros=True), dtype=dtypes.int64, device=device).reshape(1,-1)
def text_transformer_forward(self, x:Tensor, attn_mask:Optional[Tensor]=None):
for r in self.model.transformer.resblocks:
@@ -466,7 +449,7 @@ class OpenClipEncoder:
x = x + self.positional_embedding
x = self.transformer(x, attn_mask=self.attn_mask)
x = self.ln_final(x)
x = x[Tensor.arange(x.shape[0], device=x.device), tokens.argmax(axis=-1)]
x = x[:, tokens.argmax(axis=-1)]
x = x @ self.text_projection
return x
+27 -35
View File
@@ -1,24 +1,21 @@
from tinygrad import Tensor, dtypes, nn
from tinygrad import Tensor, dtypes
from tinygrad.nn import Linear, Conv2d, GroupNorm, LayerNorm
from tinygrad.device import is_dtype_supported
from typing import Optional, Union, List, Any, Tuple, Callable
from typing import Optional, Union, List, Any, Tuple
import math
# allow for monkeypatching
Linear, Conv2d, GroupNorm, LayerNorm = nn.Linear, nn.Conv2d, nn.GroupNorm, nn.LayerNorm
attention, gelu, mixed_precision_dtype = Tensor.scaled_dot_product_attention, Tensor.gelu, dtypes.float16
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/diffusionmodules/util.py#L207
def timestep_embedding(timesteps:Tensor, dim:int, max_period=10000):
half = dim // 2
freqs = (-math.log(max_period) * Tensor.arange(half, device=timesteps.device) / half).exp()
args = timesteps.unsqueeze(1) * freqs.unsqueeze(0)
out = Tensor.cat(args.cos(), args.sin(), dim=-1)
return out.cast(mixed_precision_dtype) if is_dtype_supported(mixed_precision_dtype) else out
return out.cast(dtypes.float16) if is_dtype_supported(dtypes.float16) else out
class ResBlock:
def __init__(self, channels:int, emb_channels:int, out_channels:int, num_groups:int=32):
def __init__(self, channels:int, emb_channels:int, out_channels:int):
self.in_layers = [
GroupNorm(num_groups, channels),
GroupNorm(32, channels),
Tensor.silu,
Conv2d(channels, out_channels, 3, padding=1),
]
@@ -27,7 +24,7 @@ class ResBlock:
Linear(emb_channels, out_channels),
]
self.out_layers = [
GroupNorm(num_groups, out_channels),
GroupNorm(32, out_channels),
Tensor.silu,
lambda x: x, # needed for weights loading code to work
Conv2d(out_channels, out_channels, 3, padding=1),
@@ -48,37 +45,35 @@ class CrossAttention:
self.to_v = Linear(ctx_dim, n_heads*d_head, bias=False)
self.num_heads = n_heads
self.head_size = d_head
self.attn = attention
self.to_out = [Linear(n_heads*d_head, query_dim)]
def __call__(self, x:Tensor, ctx:Optional[Tensor]=None) -> Tensor:
ctx = x if ctx is None else ctx
q,k,v = self.to_q(x), self.to_k(ctx), self.to_v(ctx)
q,k,v = [y.reshape(x.shape[0], -1, self.num_heads, self.head_size).transpose(1,2) for y in (q,k,v)]
attention = self.attn(q, k, v).transpose(1,2)
attention = Tensor.scaled_dot_product_attention(q, k, v).transpose(1,2)
h_ = attention.reshape(x.shape[0], -1, self.num_heads * self.head_size)
return h_.sequential(self.to_out)
class GEGLU:
def __init__(self, dim_in:int, dim_out:int):
self.proj = Linear(dim_in, dim_out * 2)
self.gelu = gelu
self.dim_out = dim_out
def __call__(self, x:Tensor) -> Tensor:
x, gate = self.proj(x).chunk(2, dim=-1)
return x * self.gelu(gate)
return x * gate.gelu()
class FeedForward:
def __init__(self, dim:int, mult:int=4):
self.net: tuple[GEGLU, Callable, nn.Linear] = (
self.net = [
GEGLU(dim, dim*mult),
lambda x: x, # needed for weights loading code to work
Linear(dim*mult, dim)
)
]
def __call__(self, x:Tensor) -> Tensor:
return x.sequential(list(self.net))
return x.sequential(self.net)
class BasicTransformerBlock:
def __init__(self, dim:int, ctx_dim:int, n_heads:int, d_head:int):
@@ -97,13 +92,12 @@ class BasicTransformerBlock:
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/attention.py#L619
class SpatialTransformer:
def __init__(self, channels:int, n_heads:int, d_head:int, ctx_dim:Union[int,List[int]], use_linear:bool, depth:int=1,
norm_eps:float=1e-5):
def __init__(self, channels:int, n_heads:int, d_head:int, ctx_dim:Union[int,List[int]], use_linear:bool, depth:int=1):
if isinstance(ctx_dim, int):
ctx_dim = [ctx_dim]*depth
else:
assert isinstance(ctx_dim, list) and depth == len(ctx_dim)
self.norm = GroupNorm(32, channels, eps=norm_eps)
self.norm = GroupNorm(32, channels)
assert channels == n_heads * d_head
self.proj_in = Linear(channels, channels) if use_linear else Conv2d(channels, channels, 1)
self.transformer_blocks = [BasicTransformerBlock(channels, ctx_dim[d], n_heads, d_head) for d in range(depth)]
@@ -140,9 +134,7 @@ class Upsample:
# https://github.com/Stability-AI/generative-models/blob/fbdc58cab9f4ee2be7a5e1f2e2787ecd9311942f/sgm/modules/diffusionmodules/openaimodel.py#L472
class UNetModel:
def __init__(self, adm_in_ch:Optional[int], in_ch:int, out_ch:int, model_ch:int, attention_resolutions:List[int], num_res_blocks:int,
channel_mult:List[int], transformer_depth:List[int], ctx_dim:Union[int,List[int]], use_linear:bool=False, d_head:Optional[int]=None,
n_heads:Optional[int]=None, num_groups:int=32, st_norm_eps:float=1e-5):
def __init__(self, adm_in_ch:Optional[int], in_ch:int, out_ch:int, model_ch:int, attention_resolutions:List[int], num_res_blocks:int, channel_mult:List[int], transformer_depth:List[int], ctx_dim:Union[int,List[int]], use_linear:bool=False, d_head:Optional[int]=None, n_heads:Optional[int]=None):
self.model_ch = model_ch
self.num_res_blocks = [num_res_blocks] * len(channel_mult)
@@ -182,12 +174,12 @@ class UNetModel:
for idx, mult in enumerate(channel_mult):
for _ in range(self.num_res_blocks[idx]):
layers: List[Any] = [
ResBlock(ch, time_embed_dim, model_ch*mult, num_groups),
ResBlock(ch, time_embed_dim, model_ch*mult),
]
ch = mult * model_ch
if ds in attention_resolutions:
d_head, n_heads = get_d_and_n_heads(ch)
layers.append(SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[idx], norm_eps=st_norm_eps))
layers.append(SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[idx]))
self.input_blocks.append(layers)
input_block_channels.append(ch)
@@ -201,9 +193,9 @@ class UNetModel:
d_head, n_heads = get_d_and_n_heads(ch)
self.middle_block: List = [
ResBlock(ch, time_embed_dim, ch, num_groups),
SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[-1], norm_eps=st_norm_eps),
ResBlock(ch, time_embed_dim, ch, num_groups),
ResBlock(ch, time_embed_dim, ch),
SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[-1]),
ResBlock(ch, time_embed_dim, ch),
]
self.output_blocks = []
@@ -211,13 +203,13 @@ class UNetModel:
for i in range(self.num_res_blocks[idx] + 1):
ich = input_block_channels.pop()
layers = [
ResBlock(ch + ich, time_embed_dim, model_ch*mult, num_groups),
ResBlock(ch + ich, time_embed_dim, model_ch*mult),
]
ch = model_ch * mult
if ds in attention_resolutions:
d_head, n_heads = get_d_and_n_heads(ch)
layers.append(SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[idx], norm_eps=st_norm_eps))
layers.append(SpatialTransformer(ch, n_heads, d_head, ctx_dim, use_linear, depth=transformer_depth[idx]))
if idx > 0 and i == self.num_res_blocks[idx]:
layers.append(Upsample(ch))
@@ -225,7 +217,7 @@ class UNetModel:
self.output_blocks.append(layers)
self.out = [
GroupNorm(num_groups, ch),
GroupNorm(32, ch),
Tensor.silu,
Conv2d(model_ch, out_ch, 3, padding=1),
]
@@ -238,10 +230,10 @@ class UNetModel:
assert y.shape[0] == x.shape[0]
emb = emb + y.sequential(self.label_emb[0])
if is_dtype_supported(mixed_precision_dtype):
emb = emb.cast(mixed_precision_dtype)
ctx = ctx.cast(mixed_precision_dtype)
x = x .cast(mixed_precision_dtype)
if is_dtype_supported(dtypes.float16):
emb = emb.cast(dtypes.float16)
ctx = ctx.cast(dtypes.float16)
x = x .cast(dtypes.float16)
def run(x:Tensor, bb) -> Tensor:
if isinstance(bb, ResBlock): x = bb(x, emb)
+1 -1
View File
@@ -1,6 +1,6 @@
from tinygrad import Tensor
from tinygrad.tensor import _to_np_dtype
from tinygrad.nn.onnx import OnnxRunner, OnnxValue
from tinygrad.frontend.onnx import OnnxRunner, OnnxValue
import numpy as np
import onnxruntime as ort
+1 -1
View File
@@ -81,7 +81,7 @@ def lin_to_feats(lin:Kernel, use_sts=True):
ret = [float(x) for x in ret]
if use_sts:
my_sts = dedup([(x.shape == lin.full_shape, x.is_expanded(), any(v.mask is not None for v in x.views), len(x.views)) for x in lin.sts])
my_sts = dedup([(x.shape == lin.full_shape, x.real_strides(), any(v.mask is not None for v in x.views), len(x.views)) for x in lin.sts])
assert len(my_sts) < MAX_BUFS
sts_len = 3 + 5*MAX_DIMS
for s in my_sts:
+1 -1
View File
@@ -50,7 +50,7 @@ class TestBeamSearch(unittest.TestCase):
def test_variable_shrink_prime_number(self):
v = Variable("v", 1, 400).bind(367)
a = rand(400, 367)
b = (a.shrink(((0,v), None))+1)[:367,:367].realize()
b = (a.shrink(((0,v), None))+1).reshape(367,367).realize()
np.testing.assert_allclose(b.numpy(), a.numpy()[:367]+1, atol=1e-4, rtol=1e-4)
def test_no_mutate_rawbuffers(self):
+1 -11
View File
@@ -930,7 +930,7 @@ impl<'a> Thread<'a> {
let op = ((instr >> 16) & 0x3ff) as u32;
match op {
764 | 765 | 288 | 289 | 290 | 766 | 767 | 768 | 769 => {
764 | 765 | 288 | 289 | 290 | 766 | 768 | 769 => {
let vdst = (instr & 0xff) as usize;
let sdst = ((instr >> 8) & 0x7f) as usize;
let f = |i: u32| -> usize { ((instr >> i) & 0x1ff) as usize };
@@ -944,16 +944,6 @@ impl<'a> Thread<'a> {
assert_eq!(clmp, 0);
let vcc = match op {
767 => {
let (s0, s1, s2): (u32, u32, u64) = (self.val(s0), self.val(s1), self.val(s2));
let (mul_result, overflow_mul) = (s0 as i64).overflowing_mul(s1 as i64);
let (ret, overflow_add) = mul_result.overflowing_add(s2 as i64);
let overflowed = overflow_mul || overflow_add;
if self.exec.read() {
self.vec_reg.write64(vdst, ret as u64);
}
overflowed
},
766 => {
let (s0, s1, s2): (u32, u32, u64) = (self.val(s0), self.val(s1), self.val(s2));
let (mul_result, overflow_mul) = (s0 as u64).overflowing_mul(s1 as u64);
+1 -1
View File
@@ -4,7 +4,7 @@
Only supported on 7900XTX, requires either AM (`rmmod amdgpu`) or disabling power gating on AMD (`ppfeaturemask=0xffff3fff`, don't forget to rebuild initramfs)
SQTT is implemented on top of normal tinygrad profiling, `VIZ=1 SQTT=1` to get profile pickle with sqtt data embedded in it.
SQTT is implemented on top of normal tinygrad PROFILE=1, `PROFILE=1 SQTT=1` to get profile pickle with sqtt data embedded in it.
`SQTT_BUFFER_SIZE=X` to change size of SQTT buffer (per shader engine, 6 SEs on 7900xtx) in megabytes, default 256.
-68
View File
@@ -1,68 +0,0 @@
import ctypes
from dataclasses import dataclass
import tinygrad.runtime.autogen.comgr as comgr
from tinygrad.runtime.support.compiler_amd import check
@dataclass
class InstrCtx:
pc:int=0
inst:str=""
@comgr.amd_comgr_create_disassembly_info.argtypes[2]
def instr_cb(text, user_data):
c = ctypes.cast(user_data, ctypes.POINTER(ctypes.py_object)).contents.value
c.inst = ctypes.string_at(text).decode("utf-8","replace").strip()
return comgr.AMD_COMGR_STATUS_SUCCESS
# nop callback
@comgr.amd_comgr_create_disassembly_info.argtypes[3]
def addr_cb(*args): return comgr.AMD_COMGR_STATUS_SUCCESS
def comgr_get_address_table(lib:bytes) -> dict[int, tuple[str, int]]:
check(comgr.amd_comgr_create_data(comgr.AMD_COMGR_DATA_KIND_EXECUTABLE, ctypes.byref(data_src:=comgr.amd_comgr_data_t())))
lib_buf = ctypes.create_string_buffer(lib, len(lib))
check(comgr.amd_comgr_set_data(data_src, len(lib), lib_buf))
check(comgr.amd_comgr_get_data_isa_name(data_src, isa_sz:=ctypes.c_size_t(128), isa:=(ctypes.c_char*isa_sz.value)()))
@comgr.amd_comgr_create_disassembly_info.argtypes[1]
def memory_cb(from_addr, to, size, _):
base, buf_len = ctypes.addressof(lib_buf), len(lib_buf)
start = int(from_addr) - base
if start < 0 or start >= buf_len: return 0
ctypes.memmove(to, base + start, n:=min(int(size), buf_len - start))
return n
info_src = comgr.amd_comgr_disassembly_info_t()
check(comgr.amd_comgr_create_disassembly_info(ctypes.cast(isa, ctypes.POINTER(ctypes.c_char)), memory_cb, instr_cb, addr_cb, info_src))
@comgr.amd_comgr_iterate_symbols.argtypes[1]
def sym_callback(sym, udata):
check(comgr.amd_comgr_symbol_get_info(sym, comgr.AMD_COMGR_SYMBOL_INFO_TYPE, ctypes.byref(sym_type:=ctypes.c_int())))
if sym_type.value != comgr.AMD_COMGR_SYMBOL_TYPE_FUNC: return comgr.AMD_COMGR_STATUS_SUCCESS
check(comgr.amd_comgr_symbol_get_info(sym, comgr.AMD_COMGR_SYMBOL_INFO_VALUE, ctypes.byref(vaddr:=ctypes.c_uint64())))
check(comgr.amd_comgr_symbol_get_info(sym, comgr.AMD_COMGR_SYMBOL_INFO_SIZE, ctypes.byref(size:=ctypes.c_uint64())))
check(comgr.amd_comgr_map_elf_virtual_address_to_code_object_offset(data_src, vaddr.value, ctypes.byref(offset:=ctypes.c_uint64()),
ctypes.byref(ctypes.c_uint64()), ctypes.byref(nobits:=ctypes.c_bool())))
check(nobits.value)
base = ctypes.addressof(lib_buf)
pc = base + offset.value
end = pc + size.value
addr_table = ctypes.cast(udata, ctypes.POINTER(ctypes.py_object)).contents.value
instr_ref = ctypes.py_object(ctx:=InstrCtx())
instr_ptr = ctypes.cast(ctypes.pointer(instr_ref), ctypes.c_void_p)
while pc < end:
size_read = ctypes.c_uint64(0)
ctx.pc = pc
st = comgr.amd_comgr_disassemble_instruction(info_src, ctypes.c_uint64(pc), instr_ptr, ctypes.byref(size_read))
if st == comgr.AMD_COMGR_STATUS_SUCCESS and size_read.value:
rel = (pc - base) - offset.value
addr_table[vaddr.value + rel] = (ctx.inst, int(size_read.value))
pc += size_read.value
else: # don't inf loop if comgr fails
b = ctypes.c_ubyte.from_buffer(lib_buf, pc - base).value
addr_table[vaddr.value + (pc - base - offset.value)] = (f"DISASSEMBLER ISSUE 0x{b:02x}", 1)
pc += 1
return comgr.AMD_COMGR_STATUS_SUCCESS
addr_table:dict[int, tuple[str, int]] = {}
check(comgr.amd_comgr_iterate_symbols(data_src, sym_callback, ctypes.cast(ctypes.pointer(ctypes.py_object(addr_table)), ctypes.c_void_p)))
return addr_table
+6 -7
View File
@@ -155,7 +155,6 @@ class RGP:
device_event = device_events[device]
sqtt_events = [x for x in profile if isinstance(x, ProfileSQTTEvent) and x.device == device_event.device]
if len(sqtt_events) == 0: raise RuntimeError(f"Device {device_event.device} doesn't contain SQTT data")
device_props = sqtt_events[0].props
sqtt_itrace_enabled = any([event.itrace for event in sqtt_events])
sqtt_itrace_masked = not all_same([event.itrace for event in sqtt_events])
sqtt_itrace_se_mask = functools.reduce(lambda a,b: a|b, [int(event.itrace) << event.se for event in sqtt_events], 0) if sqtt_itrace_masked else 0
@@ -193,14 +192,14 @@ class RGP:
flags=0,
trace_shader_core_clock=0x93f05080,
trace_memory_clock=0x4a723a40,
device_id={110000: 0x744c, 110003: 0x7480}[device_props['gfx_target_version']],
device_id=0x744c,
device_revision_id=0xc8,
vgprs_per_simd=1536,
sgprs_per_simd=128*16,
shader_engines=device_props['array_count'] // device_props['simd_arrays_per_engine'],
compute_unit_per_shader_engine=device_props['simd_count'] // device_props['simd_per_cu'] // (device_props['array_count'] // device_props['simd_arrays_per_engine']),
simd_per_compute_unit=device_props['simd_per_cu'],
wavefronts_per_simd=device_props['max_waves_per_simd'],
shader_engines=6,
compute_unit_per_shader_engine=16,
simd_per_compute_unit=2,
wavefronts_per_simd=16,
minimum_vgpr_alloc=4,
vgpr_alloc_granularity=8,
minimum_sgpr_alloc=128,
@@ -219,7 +218,7 @@ class RGP:
vram_bus_width=384, # 384-bit
l2_cache_size=6 * 1024 * 1024, # 6 MB
l1_cache_size=32 * 1024, # 32 KB per SIMD (?)
lds_size=device_props['lds_size_in_kb'] * 1024,
lds_size=65536, # 64 KB per CU
gpu_name=b'NAVI31',
alu_per_clock=0,
texture_per_clock=0,
-95
View File
@@ -1,95 +0,0 @@
import ctypes, pathlib, argparse, pickle, re, functools, dataclasses
from extra.sqtt.rocprof import rocprof
from extra.sqtt.disasm import comgr_get_address_table
from tinygrad.helpers import temp, DEBUG
from tinygrad.device import ProfileEvent, ProfileProgramEvent
from tinygrad.runtime.ops_amd import ProfileSQTTEvent
@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
class _ROCParseCtx:
def __init__(self, sqtt_evs:list[ProfileSQTTEvent], prog_evs:list[ProfileProgramEvent]):
self.sqtt_evs, self.prog_evs = iter(sqtt_evs), prog_evs
self.wave_events = {}
def next_sqtt(self): return next(self.sqtt_evs, None)
def find_program(self, idx): return self.prog_evs[idx]
def get_instr_info(self, idx, exec_addr): return self.disasm_program(idx)[exec_addr - self.find_program(idx).base]
@functools.lru_cache(None)
def disasm_program(self, idx): return comgr_get_address_table(self.find_program(idx).lib)
def on_occupancy_ev(self, ev):
if DEBUG >= 4: print("OCC", ev.time, ev.cu, ev.simd, ev.wave_id, ev.start)
def on_wave_ev(self, ev):
if DEBUG >= 4: print("WAVE", ev.wave_id, ev.cu, ev.simd, ev.contexts, ev.begin_time, ev.end_time)
asm = {}
for j in range(ev.instructions_size):
inst_ev = ev.instructions_array[j]
inst_typ = rocprof.rocprofiler_thread_trace_decoder_inst_category_t__enumvalues[inst_ev.category]
asm.setdefault(inst_ev.pc.address, InstInfo(typ=inst_typ, inst=self.get_instr_info(inst_ev.pc.code_object_id, inst_ev.pc.address)[0]))
asm[inst_ev.pc.address].on_ev(inst_ev)
self.wave_events[(self.find_program(ev.instructions_array[0].pc.code_object_id).name, ev.wave_id, ev.cu, ev.simd)] = asm
if __name__ == "__main__":
parser = argparse.ArgumentParser()
parser.add_argument('--profile', type=pathlib.Path, help='Path to profile', default=pathlib.Path(temp("profile.pkl", append_user=True)))
args = parser.parse_args()
with args.profile.open("rb") as f: profile = pickle.load(f)
sqtt_events:list[ProfileSQTTEvent] = []
prog_events:list[ProfileProgramEvent] = []
for e in profile:
if isinstance(e, ProfileSQTTEvent): sqtt_events.append(e)
if isinstance(e, ProfileProgramEvent) and e.device.startswith("AMD"): prog_events.append(e)
ROCParseCtx = _ROCParseCtx(sqtt_events, prog_events)
@rocprof.rocprof_trace_decoder_se_data_callback_t
def copy_cb(buf, buf_size, data_ptr):
if (prof:=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_size[0] = len(prof.blob)
return len(prof.blob)
@rocprof.rocprof_trace_decoder_trace_callback_t
def trace_cb(record_type, events_ptr, n, data_ptr):
match record_type:
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)
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)
case _:
if DEBUG >= 2: print(rocprof.rocprofiler_thread_trace_decoder_record_type_t__enumvalues[record_type], events_ptr, n)
return rocprof.ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS
@rocprof.rocprof_trace_decoder_isa_callback_t
def isa_cb(instr_ptr, mem_size_ptr, size_ptr, pc, data_ptr):
instr, mem_size_ptr[0] = ROCParseCtx.get_instr_info(pc.code_object_id, pc.address)
# this is the number of bytes to next instruction, set to 0 for end_pgm
if instr == "s_endpgm": mem_size_ptr[0] = 0
if (max_sz:=size_ptr[0]) == 0: return rocprof.ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_OUT_OF_RESOURCES
# truncate the instr if it doesn't fit
if (str_sz:=len(instr_bytes:=instr.encode()))+1 > max_sz: str_sz = max_sz
ctypes.memmove(instr_ptr, instr_bytes, str_sz)
size_ptr[0] = str_sz
return rocprof.ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS
rocprof.rocprof_trace_decoder_parse_data(copy_cb, trace_cb, isa_cb, None)
print(ROCParseCtx.wave_events.keys())
-656
View File
@@ -1,656 +0,0 @@
# pylint: skip-file
# mypy: ignore-errors
# -*- coding: utf-8 -*-
#
# TARGET arch is: []
# WORD_SIZE is: 8
# POINTER_SIZE is: 8
# LONGDOUBLE_SIZE is: 16
#
import ctypes, tinygrad.helpers.fetch as tgfetch
class AsDictMixin:
@classmethod
def as_dict(cls, self):
result = {}
if not isinstance(self, AsDictMixin):
# not a structure, assume it's already a python object
return self
if not hasattr(cls, "_fields_"):
return result
# sys.version_info >= (3, 5)
# for (field, *_) in cls._fields_: # noqa
for field_tuple in cls._fields_: # noqa
field = field_tuple[0]
if field.startswith('PADDING_'):
continue
value = getattr(self, field)
type_ = type(value)
if hasattr(value, "_length_") and hasattr(value, "_type_"):
# array
if not hasattr(type_, "as_dict"):
value = [v for v in value]
else:
type_ = type_._type_
value = [type_.as_dict(v) for v in value]
elif hasattr(value, "contents") and hasattr(value, "_type_"):
# pointer
try:
if not hasattr(type_, "as_dict"):
value = value.contents
else:
type_ = type_._type_
value = type_.as_dict(value.contents)
except ValueError:
# nullptr
value = None
elif isinstance(value, AsDictMixin):
# other structure
value = type_.as_dict(value)
result[field] = value
return result
class Structure(ctypes.Structure, AsDictMixin):
def __init__(self, *args, **kwds):
# We don't want to use positional arguments fill PADDING_* fields
args = dict(zip(self.__class__._field_names_(), args))
args.update(kwds)
super(Structure, self).__init__(**args)
@classmethod
def _field_names_(cls):
if hasattr(cls, '_fields_'):
return (f[0] for f in cls._fields_ if not f[0].startswith('PADDING'))
else:
return ()
@classmethod
def get_type(cls, field):
for f in cls._fields_:
if f[0] == field:
return f[1]
return None
@classmethod
def bind(cls, bound_fields):
fields = {}
for name, type_ in cls._fields_:
if hasattr(type_, "restype"):
if name in bound_fields:
if bound_fields[name] is None:
fields[name] = type_()
else:
# use a closure to capture the callback from the loop scope
fields[name] = (
type_((lambda callback: lambda *args: callback(*args))(
bound_fields[name]))
)
del bound_fields[name]
else:
# default callback implementation (does nothing)
try:
default_ = type_(0).restype().value
except TypeError:
default_ = None
fields[name] = type_((
lambda default_: lambda *args: default_)(default_))
else:
# not a callback function, use default initialization
if name in bound_fields:
fields[name] = bound_fields[name]
del bound_fields[name]
else:
fields[name] = type_()
if len(bound_fields) != 0:
raise ValueError(
"Cannot bind the following unknown callback(s) {}.{}".format(
cls.__name__, bound_fields.keys()
))
return cls(**fields)
class Union(ctypes.Union, AsDictMixin):
pass
c_int128 = ctypes.c_ubyte*16
c_uint128 = c_int128
void = None
if ctypes.sizeof(ctypes.c_longdouble) == 16:
c_long_double_t = ctypes.c_longdouble
else:
c_long_double_t = ctypes.c_ubyte*16
def string_cast(char_pointer, encoding='utf-8', errors='strict'):
value = ctypes.cast(char_pointer, ctypes.c_char_p).value
if value is not None and encoding is not None:
value = value.decode(encoding, errors=errors)
return value
def char_pointer_cast(string, encoding='utf-8'):
if encoding is not None:
try:
string = string.encode(encoding)
except AttributeError:
# In Python3, bytes has no encode attribute
pass
string = ctypes.c_char_p(string)
return ctypes.cast(string, ctypes.POINTER(ctypes.c_char))
class FunctionFactoryStub:
def __getattr__(self, _):
return ctypes.CFUNCTYPE(lambda y:y)
# libraries['FIXME_STUB'] explanation
# As you did not list (-l libraryname.so) a library that exports this function
# This is a non-working stub instead.
# You can either re-run clan2py with -l /path/to/library.so
# Or manually fix this by comment the ctypes.CDLL loading
_libraries = {}
_libraries['FIXME_STUB'] = ctypes.CDLL(str(tgfetch('https://github.com/ROCm/rocprof-trace-decoder/raw/5420409ad0963b2d76450add067b9058493ccbd0/releases/linux_glibc_2_28_x86_64/librocprof-trace-decoder.so'))) # ctypes.CDLL('FIXME_STUB')
# values for enumeration 'rocprofiler_thread_trace_decoder_info_t'
rocprofiler_thread_trace_decoder_info_t__enumvalues = {
0: 'ROCPROFILER_THREAD_TRACE_DECODER_INFO_NONE',
1: 'ROCPROFILER_THREAD_TRACE_DECODER_INFO_DATA_LOST',
2: 'ROCPROFILER_THREAD_TRACE_DECODER_INFO_STITCH_INCOMPLETE',
3: 'ROCPROFILER_THREAD_TRACE_DECODER_INFO_WAVE_INCOMPLETE',
4: 'ROCPROFILER_THREAD_TRACE_DECODER_INFO_LAST',
}
ROCPROFILER_THREAD_TRACE_DECODER_INFO_NONE = 0
ROCPROFILER_THREAD_TRACE_DECODER_INFO_DATA_LOST = 1
ROCPROFILER_THREAD_TRACE_DECODER_INFO_STITCH_INCOMPLETE = 2
ROCPROFILER_THREAD_TRACE_DECODER_INFO_WAVE_INCOMPLETE = 3
ROCPROFILER_THREAD_TRACE_DECODER_INFO_LAST = 4
rocprofiler_thread_trace_decoder_info_t = ctypes.c_uint32 # enum
class struct_rocprofiler_thread_trace_decoder_pc_t(Structure):
pass
struct_rocprofiler_thread_trace_decoder_pc_t._pack_ = 1 # source:False
struct_rocprofiler_thread_trace_decoder_pc_t._fields_ = [
('address', ctypes.c_uint64),
('code_object_id', ctypes.c_uint64),
]
rocprofiler_thread_trace_decoder_pc_t = struct_rocprofiler_thread_trace_decoder_pc_t
class struct_rocprofiler_thread_trace_decoder_perfevent_t(Structure):
pass
struct_rocprofiler_thread_trace_decoder_perfevent_t._pack_ = 1 # source:False
struct_rocprofiler_thread_trace_decoder_perfevent_t._fields_ = [
('time', ctypes.c_int64),
('events0', ctypes.c_uint16),
('events1', ctypes.c_uint16),
('events2', ctypes.c_uint16),
('events3', ctypes.c_uint16),
('CU', ctypes.c_ubyte),
('bank', ctypes.c_ubyte),
('PADDING_0', ctypes.c_ubyte * 6),
]
rocprofiler_thread_trace_decoder_perfevent_t = struct_rocprofiler_thread_trace_decoder_perfevent_t
class struct_rocprofiler_thread_trace_decoder_occupancy_t(Structure):
pass
struct_rocprofiler_thread_trace_decoder_occupancy_t._pack_ = 1 # source:False
struct_rocprofiler_thread_trace_decoder_occupancy_t._fields_ = [
('pc', rocprofiler_thread_trace_decoder_pc_t),
('time', ctypes.c_uint64),
('reserved', ctypes.c_ubyte),
('cu', ctypes.c_ubyte),
('simd', ctypes.c_ubyte),
('wave_id', ctypes.c_ubyte),
('start', ctypes.c_uint32, 1),
('_rsvd', ctypes.c_uint32, 31),
]
rocprofiler_thread_trace_decoder_occupancy_t = struct_rocprofiler_thread_trace_decoder_occupancy_t
# values for enumeration 'rocprofiler_thread_trace_decoder_wstate_type_t'
rocprofiler_thread_trace_decoder_wstate_type_t__enumvalues = {
0: 'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_EMPTY',
1: 'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_IDLE',
2: 'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_EXEC',
3: 'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_WAIT',
4: 'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_STALL',
5: 'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_LAST',
}
ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_EMPTY = 0
ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_IDLE = 1
ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_EXEC = 2
ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_WAIT = 3
ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_STALL = 4
ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_LAST = 5
rocprofiler_thread_trace_decoder_wstate_type_t = ctypes.c_uint32 # enum
class struct_rocprofiler_thread_trace_decoder_wave_state_t(Structure):
pass
struct_rocprofiler_thread_trace_decoder_wave_state_t._pack_ = 1 # source:False
struct_rocprofiler_thread_trace_decoder_wave_state_t._fields_ = [
('type', ctypes.c_int32),
('duration', ctypes.c_int32),
]
rocprofiler_thread_trace_decoder_wave_state_t = struct_rocprofiler_thread_trace_decoder_wave_state_t
# values for enumeration 'rocprofiler_thread_trace_decoder_inst_category_t'
rocprofiler_thread_trace_decoder_inst_category_t__enumvalues = {
0: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_NONE',
1: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_SMEM',
2: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_SALU',
3: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_VMEM',
4: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_FLAT',
5: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_LDS',
6: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_VALU',
7: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_JUMP',
8: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_NEXT',
9: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_IMMED',
10: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_CONTEXT',
11: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_MESSAGE',
12: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_BVH',
13: 'ROCPROFILER_THREAD_TRACE_DECODER_INST_LAST',
}
ROCPROFILER_THREAD_TRACE_DECODER_INST_NONE = 0
ROCPROFILER_THREAD_TRACE_DECODER_INST_SMEM = 1
ROCPROFILER_THREAD_TRACE_DECODER_INST_SALU = 2
ROCPROFILER_THREAD_TRACE_DECODER_INST_VMEM = 3
ROCPROFILER_THREAD_TRACE_DECODER_INST_FLAT = 4
ROCPROFILER_THREAD_TRACE_DECODER_INST_LDS = 5
ROCPROFILER_THREAD_TRACE_DECODER_INST_VALU = 6
ROCPROFILER_THREAD_TRACE_DECODER_INST_JUMP = 7
ROCPROFILER_THREAD_TRACE_DECODER_INST_NEXT = 8
ROCPROFILER_THREAD_TRACE_DECODER_INST_IMMED = 9
ROCPROFILER_THREAD_TRACE_DECODER_INST_CONTEXT = 10
ROCPROFILER_THREAD_TRACE_DECODER_INST_MESSAGE = 11
ROCPROFILER_THREAD_TRACE_DECODER_INST_BVH = 12
ROCPROFILER_THREAD_TRACE_DECODER_INST_LAST = 13
rocprofiler_thread_trace_decoder_inst_category_t = ctypes.c_uint32 # enum
class struct_rocprofiler_thread_trace_decoder_inst_t(Structure):
pass
struct_rocprofiler_thread_trace_decoder_inst_t._pack_ = 1 # source:False
struct_rocprofiler_thread_trace_decoder_inst_t._fields_ = [
('category', ctypes.c_uint32, 8),
('stall', ctypes.c_uint32, 24),
('duration', ctypes.c_int32),
('time', ctypes.c_int64),
('pc', rocprofiler_thread_trace_decoder_pc_t),
]
rocprofiler_thread_trace_decoder_inst_t = struct_rocprofiler_thread_trace_decoder_inst_t
class struct_rocprofiler_thread_trace_decoder_wave_t(Structure):
pass
struct_rocprofiler_thread_trace_decoder_wave_t._pack_ = 1 # source:False
struct_rocprofiler_thread_trace_decoder_wave_t._fields_ = [
('cu', ctypes.c_ubyte),
('simd', ctypes.c_ubyte),
('wave_id', ctypes.c_ubyte),
('contexts', ctypes.c_ubyte),
('_rsvd1', ctypes.c_uint32),
('_rsvd2', ctypes.c_uint32),
('_rsvd3', ctypes.c_uint32),
('begin_time', ctypes.c_int64),
('end_time', ctypes.c_int64),
('timeline_size', ctypes.c_uint64),
('instructions_size', ctypes.c_uint64),
('timeline_array', ctypes.POINTER(struct_rocprofiler_thread_trace_decoder_wave_state_t)),
('instructions_array', ctypes.POINTER(struct_rocprofiler_thread_trace_decoder_inst_t)),
]
rocprofiler_thread_trace_decoder_wave_t = struct_rocprofiler_thread_trace_decoder_wave_t
class struct_rocprofiler_thread_trace_decoder_realtime_t(Structure):
pass
struct_rocprofiler_thread_trace_decoder_realtime_t._pack_ = 1 # source:False
struct_rocprofiler_thread_trace_decoder_realtime_t._fields_ = [
('shader_clock', ctypes.c_int64),
('realtime_clock', ctypes.c_uint64),
('reserved', ctypes.c_uint64),
]
rocprofiler_thread_trace_decoder_realtime_t = struct_rocprofiler_thread_trace_decoder_realtime_t
# values for enumeration 'rocprofiler_thread_trace_decoder_shaderdata_flags_t'
rocprofiler_thread_trace_decoder_shaderdata_flags_t__enumvalues = {
0: 'ROCPROFILER_THREAD_TRACE_DECODER_SHADERDATA_FLAGS_IMM',
1: 'ROCPROFILER_THREAD_TRACE_DECODER_SHADERDATA_FLAGS_PRIV',
}
ROCPROFILER_THREAD_TRACE_DECODER_SHADERDATA_FLAGS_IMM = 0
ROCPROFILER_THREAD_TRACE_DECODER_SHADERDATA_FLAGS_PRIV = 1
rocprofiler_thread_trace_decoder_shaderdata_flags_t = ctypes.c_uint32 # enum
class struct_rocprofiler_thread_trace_decoder_shaderdata_t(Structure):
pass
struct_rocprofiler_thread_trace_decoder_shaderdata_t._pack_ = 1 # source:False
struct_rocprofiler_thread_trace_decoder_shaderdata_t._fields_ = [
('time', ctypes.c_int64),
('value', ctypes.c_uint64),
('cu', ctypes.c_ubyte),
('simd', ctypes.c_ubyte),
('wave_id', ctypes.c_ubyte),
('flags', ctypes.c_ubyte),
('reserved', ctypes.c_uint32),
]
rocprofiler_thread_trace_decoder_shaderdata_t = struct_rocprofiler_thread_trace_decoder_shaderdata_t
# values for enumeration 'rocprofiler_thread_trace_decoder_record_type_t'
rocprofiler_thread_trace_decoder_record_type_t__enumvalues = {
0: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_GFXIP',
1: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_OCCUPANCY',
2: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_PERFEVENT',
3: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_WAVE',
4: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_INFO',
5: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_DEBUG',
6: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_SHADERDATA',
7: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_REALTIME',
8: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_RT_FREQUENCY',
9: 'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_LAST',
}
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_GFXIP = 0
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_OCCUPANCY = 1
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_PERFEVENT = 2
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_WAVE = 3
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_INFO = 4
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_DEBUG = 5
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_SHADERDATA = 6
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_REALTIME = 7
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_RT_FREQUENCY = 8
ROCPROFILER_THREAD_TRACE_DECODER_RECORD_LAST = 9
rocprofiler_thread_trace_decoder_record_type_t = ctypes.c_uint32 # enum
# values for enumeration 'c__EA_rocprofiler_thread_trace_decoder_status_t'
c__EA_rocprofiler_thread_trace_decoder_status_t__enumvalues = {
0: 'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS',
1: 'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR',
2: 'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_OUT_OF_RESOURCES',
3: 'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_INVALID_ARGUMENT',
4: 'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_INVALID_SHADER_DATA',
5: 'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_LAST',
}
ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS = 0
ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR = 1
ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_OUT_OF_RESOURCES = 2
ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_INVALID_ARGUMENT = 3
ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_INVALID_SHADER_DATA = 4
ROCPROFILER_THREAD_TRACE_DECODER_STATUS_LAST = 5
c__EA_rocprofiler_thread_trace_decoder_status_t = ctypes.c_uint32 # enum
rocprofiler_thread_trace_decoder_status_t = c__EA_rocprofiler_thread_trace_decoder_status_t
rocprofiler_thread_trace_decoder_status_t__enumvalues = c__EA_rocprofiler_thread_trace_decoder_status_t__enumvalues
rocprof_trace_decoder_trace_callback_t = ctypes.CFUNCTYPE(c__EA_rocprofiler_thread_trace_decoder_status_t, rocprofiler_thread_trace_decoder_record_type_t, ctypes.POINTER(None), ctypes.c_uint64, ctypes.POINTER(None))
rocprof_trace_decoder_isa_callback_t = ctypes.CFUNCTYPE(c__EA_rocprofiler_thread_trace_decoder_status_t, ctypes.POINTER(ctypes.c_char), ctypes.POINTER(ctypes.c_uint64), ctypes.POINTER(ctypes.c_uint64), struct_rocprofiler_thread_trace_decoder_pc_t, ctypes.POINTER(None))
rocprof_trace_decoder_se_data_callback_t = ctypes.CFUNCTYPE(ctypes.c_uint64, ctypes.POINTER(ctypes.POINTER(ctypes.c_ubyte)), ctypes.POINTER(ctypes.c_uint64), ctypes.POINTER(None))
try:
rocprof_trace_decoder_parse_data = _libraries['FIXME_STUB'].rocprof_trace_decoder_parse_data
rocprof_trace_decoder_parse_data.restype = rocprofiler_thread_trace_decoder_status_t
rocprof_trace_decoder_parse_data.argtypes = [rocprof_trace_decoder_se_data_callback_t, rocprof_trace_decoder_trace_callback_t, rocprof_trace_decoder_isa_callback_t, ctypes.POINTER(None)]
except AttributeError:
pass
try:
rocprof_trace_decoder_get_info_string = _libraries['FIXME_STUB'].rocprof_trace_decoder_get_info_string
rocprof_trace_decoder_get_info_string.restype = ctypes.POINTER(ctypes.c_char)
rocprof_trace_decoder_get_info_string.argtypes = [rocprofiler_thread_trace_decoder_info_t]
except AttributeError:
pass
try:
rocprof_trace_decoder_get_status_string = _libraries['FIXME_STUB'].rocprof_trace_decoder_get_status_string
rocprof_trace_decoder_get_status_string.restype = ctypes.POINTER(ctypes.c_char)
rocprof_trace_decoder_get_status_string.argtypes = [rocprofiler_thread_trace_decoder_status_t]
except AttributeError:
pass
rocprofiler_thread_trace_decoder_debug_callback_t = ctypes.CFUNCTYPE(None, ctypes.c_int64, ctypes.POINTER(ctypes.c_char), ctypes.POINTER(ctypes.c_char), ctypes.POINTER(None))
uint64_t = ctypes.c_uint64
try:
rocprof_trace_decoder_dump_data = _libraries['FIXME_STUB'].rocprof_trace_decoder_dump_data
rocprof_trace_decoder_dump_data.restype = rocprofiler_thread_trace_decoder_status_t
rocprof_trace_decoder_dump_data.argtypes = [ctypes.POINTER(ctypes.c_char), uint64_t, rocprofiler_thread_trace_decoder_debug_callback_t, ctypes.POINTER(None)]
except AttributeError:
pass
class union_rocprof_trace_decoder_gfx9_header_t(Union):
pass
class struct_rocprof_trace_decoder_gfx9_header_t_0(Structure):
pass
struct_rocprof_trace_decoder_gfx9_header_t_0._pack_ = 1 # source:False
struct_rocprof_trace_decoder_gfx9_header_t_0._fields_ = [
('legacy_version', ctypes.c_uint64, 13),
('gfx9_version2', ctypes.c_uint64, 3),
('DSIMDM', ctypes.c_uint64, 4),
('DCU', ctypes.c_uint64, 5),
('reserved1', ctypes.c_uint64, 1),
('SEID', ctypes.c_uint64, 6),
('reserved2', ctypes.c_uint64, 32),
]
union_rocprof_trace_decoder_gfx9_header_t._pack_ = 1 # source:False
union_rocprof_trace_decoder_gfx9_header_t._anonymous_ = ('_0',)
union_rocprof_trace_decoder_gfx9_header_t._fields_ = [
('_0', struct_rocprof_trace_decoder_gfx9_header_t_0),
('raw', ctypes.c_uint64),
]
rocprof_trace_decoder_gfx9_header_t = union_rocprof_trace_decoder_gfx9_header_t
class union_rocprof_trace_decoder_instrument_enable_t(Union):
pass
class struct_rocprof_trace_decoder_instrument_enable_t_0(Structure):
pass
struct_rocprof_trace_decoder_instrument_enable_t_0._pack_ = 1 # source:False
struct_rocprof_trace_decoder_instrument_enable_t_0._fields_ = [
('char1', ctypes.c_uint32, 8),
('char2', ctypes.c_uint32, 8),
('char3', ctypes.c_uint32, 8),
('char4', ctypes.c_uint32, 8),
]
union_rocprof_trace_decoder_instrument_enable_t._pack_ = 1 # source:False
union_rocprof_trace_decoder_instrument_enable_t._anonymous_ = ('_0',)
union_rocprof_trace_decoder_instrument_enable_t._fields_ = [
('_0', struct_rocprof_trace_decoder_instrument_enable_t_0),
('u32All', ctypes.c_uint32),
]
rocprof_trace_decoder_instrument_enable_t = union_rocprof_trace_decoder_instrument_enable_t
class union_rocprof_trace_decoder_packet_header_t(Union):
pass
class struct_rocprof_trace_decoder_packet_header_t_0(Structure):
pass
struct_rocprof_trace_decoder_packet_header_t_0._pack_ = 1 # source:False
struct_rocprof_trace_decoder_packet_header_t_0._fields_ = [
('opcode', ctypes.c_uint32, 8),
('type', ctypes.c_uint32, 4),
('data20', ctypes.c_uint32, 20),
]
union_rocprof_trace_decoder_packet_header_t._pack_ = 1 # source:False
union_rocprof_trace_decoder_packet_header_t._anonymous_ = ('_0',)
union_rocprof_trace_decoder_packet_header_t._fields_ = [
('_0', struct_rocprof_trace_decoder_packet_header_t_0),
('u32All', ctypes.c_uint32),
]
rocprof_trace_decoder_packet_header_t = union_rocprof_trace_decoder_packet_header_t
# values for enumeration 'rocprof_trace_decoder_packet_opcode_t'
rocprof_trace_decoder_packet_opcode_t__enumvalues = {
4: 'ROCPROF_TRACE_DECODER_PACKET_OPCODE_CODEOBJ',
5: 'ROCPROF_TRACE_DECODER_PACKET_OPCODE_RT_TIMESTAMP',
6: 'ROCPROF_TRACE_DECODER_PACKET_OPCODE_AGENT_INFO',
}
ROCPROF_TRACE_DECODER_PACKET_OPCODE_CODEOBJ = 4
ROCPROF_TRACE_DECODER_PACKET_OPCODE_RT_TIMESTAMP = 5
ROCPROF_TRACE_DECODER_PACKET_OPCODE_AGENT_INFO = 6
rocprof_trace_decoder_packet_opcode_t = ctypes.c_uint32 # enum
# values for enumeration 'rocprof_trace_decoder_agent_info_type_t'
rocprof_trace_decoder_agent_info_type_t__enumvalues = {
0: 'ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_RT_FREQUENCY_KHZ',
1: 'ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_COUNTER_INTERVAL',
2: 'ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_LAST',
}
ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_RT_FREQUENCY_KHZ = 0
ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_COUNTER_INTERVAL = 1
ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_LAST = 2
rocprof_trace_decoder_agent_info_type_t = ctypes.c_uint32 # enum
class union_rocprof_trace_decoder_codeobj_marker_tail_t(Union):
pass
class struct_rocprof_trace_decoder_codeobj_marker_tail_t_0(Structure):
pass
struct_rocprof_trace_decoder_codeobj_marker_tail_t_0._pack_ = 1 # source:False
struct_rocprof_trace_decoder_codeobj_marker_tail_t_0._fields_ = [
('isUnload', ctypes.c_uint32, 1),
('bFromStart', ctypes.c_uint32, 1),
('legacy_id', ctypes.c_uint32, 30),
]
union_rocprof_trace_decoder_codeobj_marker_tail_t._pack_ = 1 # source:False
union_rocprof_trace_decoder_codeobj_marker_tail_t._anonymous_ = ('_0',)
union_rocprof_trace_decoder_codeobj_marker_tail_t._fields_ = [
('_0', struct_rocprof_trace_decoder_codeobj_marker_tail_t_0),
('raw', ctypes.c_uint32),
]
rocprof_trace_decoder_codeobj_marker_tail_t = union_rocprof_trace_decoder_codeobj_marker_tail_t
# values for enumeration 'rocprof_trace_decoder_codeobj_marker_type_t'
rocprof_trace_decoder_codeobj_marker_type_t__enumvalues = {
0: 'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_TAIL',
1: 'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_SIZE_LO',
2: 'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ADDR_LO',
3: 'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ADDR_HI',
4: 'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_SIZE_HI',
5: 'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ID_LO',
6: 'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ID_HI',
7: 'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_LAST',
}
ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_TAIL = 0
ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_SIZE_LO = 1
ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ADDR_LO = 2
ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ADDR_HI = 3
ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_SIZE_HI = 4
ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ID_LO = 5
ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ID_HI = 6
ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_LAST = 7
rocprof_trace_decoder_codeobj_marker_type_t = ctypes.c_uint32 # enum
__all__ = \
['ROCPROFILER_THREAD_TRACE_DECODER_INFO_DATA_LOST',
'ROCPROFILER_THREAD_TRACE_DECODER_INFO_LAST',
'ROCPROFILER_THREAD_TRACE_DECODER_INFO_NONE',
'ROCPROFILER_THREAD_TRACE_DECODER_INFO_STITCH_INCOMPLETE',
'ROCPROFILER_THREAD_TRACE_DECODER_INFO_WAVE_INCOMPLETE',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_BVH',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_CONTEXT',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_FLAT',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_IMMED',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_JUMP',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_LAST',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_LDS',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_MESSAGE',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_NEXT',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_NONE',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_SALU',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_SMEM',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_VALU',
'ROCPROFILER_THREAD_TRACE_DECODER_INST_VMEM',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_DEBUG',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_GFXIP',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_INFO',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_LAST',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_OCCUPANCY',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_PERFEVENT',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_REALTIME',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_RT_FREQUENCY',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_SHADERDATA',
'ROCPROFILER_THREAD_TRACE_DECODER_RECORD_WAVE',
'ROCPROFILER_THREAD_TRACE_DECODER_SHADERDATA_FLAGS_IMM',
'ROCPROFILER_THREAD_TRACE_DECODER_SHADERDATA_FLAGS_PRIV',
'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR',
'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_INVALID_ARGUMENT',
'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_INVALID_SHADER_DATA',
'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_ERROR_OUT_OF_RESOURCES',
'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_LAST',
'ROCPROFILER_THREAD_TRACE_DECODER_STATUS_SUCCESS',
'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_EMPTY',
'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_EXEC',
'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_IDLE',
'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_LAST',
'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_STALL',
'ROCPROFILER_THREAD_TRACE_DECODER_WSTATE_WAIT',
'ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_COUNTER_INTERVAL',
'ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_LAST',
'ROCPROF_TRACE_DECODER_AGENT_INFO_TYPE_RT_FREQUENCY_KHZ',
'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ADDR_HI',
'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ADDR_LO',
'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ID_HI',
'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_ID_LO',
'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_LAST',
'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_SIZE_HI',
'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_SIZE_LO',
'ROCPROF_TRACE_DECODER_CODEOBJ_MARKER_TYPE_TAIL',
'ROCPROF_TRACE_DECODER_PACKET_OPCODE_AGENT_INFO',
'ROCPROF_TRACE_DECODER_PACKET_OPCODE_CODEOBJ',
'ROCPROF_TRACE_DECODER_PACKET_OPCODE_RT_TIMESTAMP',
'c__EA_rocprofiler_thread_trace_decoder_status_t',
'rocprof_trace_decoder_agent_info_type_t',
'rocprof_trace_decoder_codeobj_marker_tail_t',
'rocprof_trace_decoder_codeobj_marker_type_t',
'rocprof_trace_decoder_dump_data',
'rocprof_trace_decoder_get_info_string',
'rocprof_trace_decoder_get_status_string',
'rocprof_trace_decoder_gfx9_header_t',
'rocprof_trace_decoder_instrument_enable_t',
'rocprof_trace_decoder_isa_callback_t',
'rocprof_trace_decoder_packet_header_t',
'rocprof_trace_decoder_packet_opcode_t',
'rocprof_trace_decoder_parse_data',
'rocprof_trace_decoder_se_data_callback_t',
'rocprof_trace_decoder_trace_callback_t',
'rocprofiler_thread_trace_decoder_debug_callback_t',
'rocprofiler_thread_trace_decoder_info_t',
'rocprofiler_thread_trace_decoder_inst_category_t',
'rocprofiler_thread_trace_decoder_inst_t',
'rocprofiler_thread_trace_decoder_occupancy_t',
'rocprofiler_thread_trace_decoder_pc_t',
'rocprofiler_thread_trace_decoder_perfevent_t',
'rocprofiler_thread_trace_decoder_realtime_t',
'rocprofiler_thread_trace_decoder_record_type_t',
'rocprofiler_thread_trace_decoder_shaderdata_flags_t',
'rocprofiler_thread_trace_decoder_shaderdata_t',
'rocprofiler_thread_trace_decoder_status_t',
'rocprofiler_thread_trace_decoder_status_t__enumvalues',
'rocprofiler_thread_trace_decoder_wave_state_t',
'rocprofiler_thread_trace_decoder_wave_t',
'rocprofiler_thread_trace_decoder_wstate_type_t',
'struct_rocprof_trace_decoder_codeobj_marker_tail_t_0',
'struct_rocprof_trace_decoder_gfx9_header_t_0',
'struct_rocprof_trace_decoder_instrument_enable_t_0',
'struct_rocprof_trace_decoder_packet_header_t_0',
'struct_rocprofiler_thread_trace_decoder_inst_t',
'struct_rocprofiler_thread_trace_decoder_occupancy_t',
'struct_rocprofiler_thread_trace_decoder_pc_t',
'struct_rocprofiler_thread_trace_decoder_perfevent_t',
'struct_rocprofiler_thread_trace_decoder_realtime_t',
'struct_rocprofiler_thread_trace_decoder_shaderdata_t',
'struct_rocprofiler_thread_trace_decoder_wave_state_t',
'struct_rocprofiler_thread_trace_decoder_wave_t', 'uint64_t',
'union_rocprof_trace_decoder_codeobj_marker_tail_t',
'union_rocprof_trace_decoder_gfx9_header_t',
'union_rocprof_trace_decoder_instrument_enable_t',
'union_rocprof_trace_decoder_packet_header_t']
+40
View File
@@ -0,0 +1,40 @@
import time
from extra.optimization.helpers import load_worlds, ast_str_to_ast
from tinygrad import Device
from tinygrad.codegen.lowerer import pm_lowerer, get_index
from tinygrad.uop.ops import graph_rewrite
from tinygrad.codegen.opt.kernel import Kernel
from tinygrad.codegen.opt.postrange import Scheduler
from tinygrad.codegen.opt.heuristic import hand_coded_optimizations
from tinygrad.helpers import getenv
if __name__ == "__main__":
renderer = Device.default.renderer
ast_strs = load_worlds()
if (n:=getenv("N", -1)) != -1: ast_strs = ast_strs[n:n+1]
good = 0
for i, ast_str in enumerate(ast_strs):
ast = ast_str_to_ast(ast_str)
st = time.perf_counter()
lin = Kernel(ast, renderer)
opt1 = hand_coded_optimizations(lin)
et_lin = time.perf_counter() - st
lowered = graph_rewrite(ast, pm_lowerer, ctx=get_index(ast), bottom_up=True)
st = time.perf_counter()
sch = Scheduler(lowered, renderer)
sch.convert_loop_to_global()
sch.simplify_merge_adjacent()
opt2 = hand_coded_optimizations(sch)
et_sch = time.perf_counter() - st
if opt1 != opt2:
print(f"******* {i:6d}")
print("Kernel: ", lin.colored_shape(), "->", lin.apply_opts(opt1).colored_shape())
print("Scheduler: ", sch.colored_shape(), "->", sch.apply_opts(opt2).colored_shape())
print(opt1)
print(opt2)
else:
good += 1
print(f"******* {i:6d} MATCH {good/(i+1)*100:.2f}% -- {et_lin/et_sch:4.2f}x speedup")
@@ -1,400 +0,0 @@
/**
* @file
* @brief Basic operations on generic types.
*/
#pragma once
#include <cuda_bf16.h>
#include <limits>
#include "base_types.cuh"
namespace kittens {
/**
* @namespace base_ops
*
* @brief A namespace for operations on basic data types.
*/
namespace base_ops {
/* ---------- CONST OPS ---------- */
/**
* @brief Represents the zero constant operation.
*
* This operation returns the zero value of the specified type.
*
* @tparam T The data type for which to return the zero value.
* @return The zero value of type T.
*/
struct zero {
template<typename T, typename... args> __device__ static inline constexpr T op(args... _) { return base_types::constants<T>::zero(); }
};
/**
* @brief Represents the one constant operation.
*
* This operation returns the one value of the specified type.
*
* @tparam T The data type for which to return the one value.
* @return The one value of type T.
*/
struct one {
template<typename T, typename... args> __device__ static inline constexpr T op(args... _) { return base_types::constants<T>::one(); }
};
/**
* @brief Represents the positive infinity constant operation.
*
* This operation returns the positive infinity value of the specified type.
*
* @tparam T The data type for which to return the positive infinity value.
* @return The positive infinity value of type T.
*/
struct pos_infty {
template<typename T, typename... args> __device__ static inline constexpr T op(args... _) { return base_types::constants<T>::pos_infty(); }
};
/**
* @brief Represents the negative infinity constant operation.
*
* This operation returns the negative infinity value of the specified type.
*
* @tparam T The data type for which to return the negative infinity value.
* @return The negative infinity value of type T.
*/
struct neg_infty {
template<typename T, typename... args> __device__ static inline constexpr T op(args... _) { return base_types::constants<T>::neg_infty(); }
};
/* ---------- UNARY OPS ---------- */
/**
* @brief Exponential function operation.
*
* This operation calculates the exponential of the input value.
*
* @tparam T The data type of the input and output values.
* @param x[in] The input value.
* @return The exponential of the input value.
*/
struct exp {
template<typename T> static __device__ inline T op(const T &x) { return exp(x); }
};
template<> __device__ inline float exp::op<float> (const float &x ) { return __expf(x); }
template<> __device__ inline float2 exp::op<float2>(const float2 &x) { return float2{__expf(x.x), __expf(x.y)}; }
template<> __device__ inline bf16 exp::op<bf16> (const bf16 &x ) { return hexp(x); }
template<> __device__ inline bf16_2 exp::op<bf16_2>(const bf16_2 &x) { return h2exp(x); }
template<> __device__ inline half exp::op<half> (const half &x ) { return hexp(x); }
template<> __device__ inline half_2 exp::op<half_2>(const half_2 &x) { return h2exp(x); }
/**
* @brief Exponential function operation, in base 2
*
* This operation calculates the exponential of the input value, in base 2.
*
* @tparam T The data type of the input and output values.
* @param x[in] The input value.
* @return The exponential of the input value.
*/
struct exp2 {
template<typename T> static __device__ inline T op(const T &x) { return exp2f(x); }
};
template<> __device__ inline float exp2::op<float> (const float &x ) { return exp2f(x); }
template<> __device__ inline float2 exp2::op<float2>(const float2 &x) { return float2{exp2f(x.x), exp2f(x.y)}; }
template<> __device__ inline bf16 exp2::op<bf16> (const bf16 &x ) { return hexp2(x); }
template<> __device__ inline bf16_2 exp2::op<bf16_2>(const bf16_2 &x) { return h2exp2(x); }
template<> __device__ inline half exp2::op<half> (const half &x ) { return hexp2(x); }
template<> __device__ inline half_2 exp2::op<half_2>(const half_2 &x) { return h2exp2(x); }
/**
* @brief Natural log function operation.
*
* This operation calculates the natural logarithm of the input value.
*
* @tparam T The data type of the input and output values.
* @param x[in] The input value.
* @return The natural logarithm of the input value.
*/
struct log {
template<typename T> static __device__ inline T op(const T &x) { return log(x); }
};
template<> __device__ inline float log::op<float> (const float &x ) { return __logf(x); }
template<> __device__ inline float2 log::op<float2>(const float2 &x) { return float2{__logf(x.x), __logf(x.y)}; }
template<> __device__ inline bf16 log::op<bf16> (const bf16 &x ) { return hlog(x); }
template<> __device__ inline bf16_2 log::op<bf16_2>(const bf16_2 &x) { return h2log(x); }
template<> __device__ inline half log::op<half> (const half &x ) { return hlog(x); }
template<> __device__ inline half_2 log::op<half_2>(const half_2 &x) { return h2log(x); }
/**
* @brief Logarithm base 2 operation.
*
* This operation calculates the logarithm base 2 of the input value.
*
* @tparam T The data type of the input and output values.
* @param x[in] The input value.
* @return The logarithm base 2 of the input value.
*/
struct log2 {
template<typename T> static __device__ inline T op(const T &x) { return log2(x); }
};
template<> __device__ inline float log2::op<float> (const float &x ) { return __log2f(x); }
template<> __device__ inline float2 log2::op<float2>(const float2 &x) { return float2{__log2f(x.x), __log2f(x.y)}; }
template<> __device__ inline bf16 log2::op<bf16> (const bf16 &x ) { return hlog2(x); }
template<> __device__ inline bf16_2 log2::op<bf16_2>(const bf16_2 &x) { return h2log2(x); }
template<> __device__ inline half log2::op<half> (const half &x ) { return hlog2(x); }
template<> __device__ inline half_2 log2::op<half_2>(const half_2 &x) { return h2log2(x); }
/**
* @brief Absolute value operation.
*
* This operation calculates the absolute value of the input.
*
* @tparam T The data type of the input and output values.
* @param x[in] The input value.
* @return The absolute value of the input.
*/
struct abs {
template<typename T> static __device__ inline T op(const T &x) { return abs(x); }
};
template<> __device__ inline float abs::op<float> (const float &x ) { return fabsf(x); }
template<> __device__ inline float2 abs::op<float2>(const float2 &x) { return float2{fabsf(x.x), fabsf(x.y)}; }
template<> __device__ inline bf16 abs::op<bf16> (const bf16 &x ) { return __habs(x); }
template<> __device__ inline bf16_2 abs::op<bf16_2>(const bf16_2 &x) { return __habs2(x); }
template<> __device__ inline half abs::op<half> (const half &x ) { return __habs(x); }
template<> __device__ inline half_2 abs::op<half_2>(const half_2 &x) { return __habs2(x); }
/**
* @brief Rectified Linear Unit (ReLU) operation.
*
* This operation applies the ReLU function to the input, which is the
* maximum of zero and the input value.
*
* @tparam T The data type of the input and output values.
* @param x[in] The input value.
* @return The result of ReLU function applied to the input.
*/
struct relu {
template<typename T> static __device__ inline T op(const T &x) { return max(x, base_types::constants<T>::zero()); }
};
template<> __device__ inline float relu::op<float> (const float &x ) { return max(x, 0.f); }
template<> __device__ inline float2 relu::op<float2>(const float2 &x) { return float2{max(x.x, 0.f), max(x.y, 0.f)}; }
template<> __device__ inline bf16 relu::op<bf16> (const bf16 &x ) { return __hmax(x, base_types::constants<bf16>::zero()); }
template<> __device__ inline bf16_2 relu::op<bf16_2>(const bf16_2 &x) { return __hmax2(x, base_types::constants<bf16_2>::zero()); }
template<> __device__ inline half relu::op<half> (const half &x ) { return __hmax(x, base_types::constants<half>::zero()); }
template<> __device__ inline half_2 relu::op<half_2>(const half_2 &x) { return __hmax2(x, base_types::constants<half_2>::zero()); }
/**
* @brief Copy operation.
*
* This operation returns the input value unchanged.
*
* @tparam T The data type of the input and output values.
* @param a[in] The input value.
* @return The same value as the input.
*/
struct copy { // for non-compile-time setters.
template<typename T> static __device__ inline T op(const T &a) { return a; }
};
/* ---------- BINARY OPS ---------- */
/**
* @brief Copy2 operation.
*
* This operation returns the second input value unchanged.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value (ignored).
* @param b[in] The second input value.
* @return The same value as the second input.
*/
struct copy2 { // this turns out to be a slightly hacky op that makes some code cleaner :/
template<typename T> static __device__ inline T op(const T &a, const T &b) { return b; }
};
/**
* @brief Sum operation.
*
* This operation calculates the sum of two input values.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value.
* @param b[in] The second input value.
* @return The sum of the input values.
*/
struct sum {
template<typename T> static __device__ inline T op(const T &a, const T &b) { return a+b; }
};
template<> __device__ inline float2 sum::op<float2>(const float2 &a, const float2 &b) {
#ifdef KITTENS_BLACKWELL
float2 c;
asm volatile("add.f32x2 %0, %1, %2;" : "=l"(*(uint64_t*)&c) : "l"(*(uint64_t*)&a), "l"(*(uint64_t*)&b));
return c;
#else
return float2{a.x+b.x, a.y+b.y};
#endif
}
template<> __device__ inline bf16 sum::op<bf16> (const bf16 &a, const bf16 &b) { return __hadd(a, b); }
template<> __device__ inline bf16_2 sum::op<bf16_2>(const bf16_2 &a, const bf16_2 &b) { return __hadd2(a, b); }
template<> __device__ inline half sum::op<half> (const half &a, const half &b) { return __hadd(a, b); }
template<> __device__ inline half_2 sum::op<half_2>(const half_2 &a, const half_2 &b) { return __hadd2(a, b); }
/**
* @brief Subtraction operation.
*
* This operation calculates the difference between two input values.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value.
* @param b[in] The second input value.
* @return The difference between the input values.
*/
struct sub {
template<typename T> static __device__ inline T op(const T &a, const T &b) { return a-b; }
};
template<> __device__ inline float2 sub::op<float2>(const float2 &a, const float2 &b) {
#ifdef KITTENS_BLACKWELL
float2 c;
asm volatile("sub.f32x2 %0, %1, %2;" : "=l"(*(uint64_t*)&c) : "l"(*(uint64_t*)&a), "l"(*(uint64_t*)&b));
return c;
#else
return float2{a.x-b.x, a.y-b.y};
#endif
}
template<> __device__ inline bf16 sub::op<bf16> (const bf16 &a, const bf16 &b) { return __hsub(a, b); }
template<> __device__ inline bf16_2 sub::op<bf16_2>(const bf16_2 &a, const bf16_2 &b) { return __hsub2(a, b); }
template<> __device__ inline half sub::op<half> (const half &a, const half &b) { return __hsub(a, b); }
template<> __device__ inline half_2 sub::op<half_2>(const half_2 &a, const half_2 &b) { return __hsub2(a, b); }
/**
* @brief Multiplication operation.
*
* This operation calculates the product of two input values.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value.
* @param b[in] The second input value.
* @return The product of the input values.
*/
struct mul {
template<typename T> static __device__ inline T op(const T &a, const T &b) { return a*b; }
};
template<> __device__ inline float2 mul::op<float2>(const float2 &a, const float2 &b) {
#ifdef KITTENS_BLACKWELL
float2 c;
asm volatile("mul.f32x2 %0, %1, %2;" : "=l"(*(uint64_t*)&c) : "l"(*(uint64_t*)&a), "l"(*(uint64_t*)&b));
return c;
#else
return float2{a.x*b.x, a.y*b.y};
#endif
}
template<> __device__ inline bf16 mul::op<bf16> (const bf16 &a, const bf16 &b) { return __hmul(a, b); }
template<> __device__ inline bf16_2 mul::op<bf16_2>(const bf16_2 &a, const bf16_2 &b) { return __hmul2(a, b); }
template<> __device__ inline half mul::op<half> (const half &a, const half &b) { return __hmul(a, b); }
template<> __device__ inline half_2 mul::op<half_2>(const half_2 &a, const half_2 &b) { return __hmul2(a, b); }
/**
* @brief Division operation.
*
* This operation calculates the quotient of two input values.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value.
* @param b[in] The second input value.
* @return The quotient of the input values.
*/
struct div {
template<typename T> static __device__ inline T op(const T &a, const T &b) { return a/b; }
};
template<> __device__ inline float2 div::op<float2>(const float2 &a, const float2 &b) { return float2{a.x/b.x, a.y/b.y}; }
template<> __device__ inline bf16 div::op<bf16> (const bf16 &a, const bf16 &b) { return __hdiv(a, b); }
template<> __device__ inline bf16_2 div::op<bf16_2>(const bf16_2 &a, const bf16_2 &b) { return __h2div(a, b); } // this op is a special snowflake
template<> __device__ inline half div::op<half> (const half &a, const half &b) { return __hdiv(a, b); }
template<> __device__ inline half_2 div::op<half_2>(const half_2 &a, const half_2 &b) { return __h2div(a, b); }
/**
* @brief Maximum operation.
*
* This operation calculates the maximum of two input values.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value.
* @param b[in] The second input value.
* @return The maximum of the input values.
*/
struct max {
template<typename T> static __device__ inline T op(const T &a, const T &b) { return ::max(a, b); }
};
template<> __device__ inline float2 max::op<float2>(const float2 &a, const float2 &b) { return float2{::max(a.x, b.x), ::max(a.y, b.y)}; }
template<> __device__ inline bf16 max::op<bf16> (const bf16 &a, const bf16 &b) { return __hmax(a, b); }
template<> __device__ inline bf16_2 max::op<bf16_2>(const bf16_2 &a, const bf16_2 &b) { return __hmax2(a, b); }
template<> __device__ inline half max::op<half> (const half &a, const half &b) { return __hmax(a, b); }
template<> __device__ inline half_2 max::op<half_2>(const half_2 &a, const half_2 &b) { return __hmax2(a, b); }
/**
* @brief Minimum operation.
*
* This operation calculates the minimum of two input values.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value.
* @param b[in] The second input value.
* @return The minimum of the input values.
*/
struct min {
template<typename T> static __device__ inline T op(const T &a, const T &b) { return ::min(a, b); }
};
template<> __device__ inline float2 min::op<float2>(const float2 &a, const float2 &b) { return float2{::min(a.x, b.x), ::min(a.y, b.y)}; }
template<> __device__ inline bf16 min::op<bf16> (const bf16 &a, const bf16 &b) { return __hmin(a, b); }
template<> __device__ inline bf16_2 min::op<bf16_2>(const bf16_2 &a, const bf16_2 &b) { return __hmin2(a, b); }
template<> __device__ inline half min::op<half> (const half &a, const half &b) { return __hmin(a, b); }
template<> __device__ inline half_2 min::op<half_2>(const half_2 &a, const half_2 &b) { return __hmin2(a, b); }
/* ---------- TERNARY OPS ---------- */
/**
* @brief Fused multiply-add operation A * B + C.
*
* This operation performs a fused multiply-add, computing (A * B) + C with only one rounding.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value.
* @param b[in] The second input value.
* @param c[in] The third input value to be added.
* @return The result of the fused multiply-add operation.
*/
struct fma_AxBtC {
template<typename T> static __device__ inline T op(const T &a, const T &b, const T &c) {
return sum::op<T>(mul::op<T>(a, b), c);
}
};
template<> __device__ inline float2 fma_AxBtC::op<float2>(const float2 &a, const float2 &b, const float2 &c) {
#ifdef KITTENS_BLACKWELL
float2 d;
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;" : "=l"(*(uint64_t*)&d) : "l"(*(uint64_t*)&a), "l"(*(uint64_t*)&b), "l"(*(uint64_t*)&c));
return d;
#else
return float2{a.x*b.x+c.x, a.y*b.y+c.y};
#endif
}
/**
* @brief Fused multiply-add operation A * C + B.
*
* This operation performs a fused multiply-add, computing (A * C) + B with only one rounding.
* This is particularly useful for attention mechanisms in neural networks.
*
* @tparam T The data type of the input and output values.
* @param a[in] The first input value.
* @param b[in] The third input value to be added.
* @param c[in] The second input value.
* @return The result of the fused multiply-add operation.
*/
struct fma_AxCtB { // this is the one needed for attention
template<typename T> static __device__ inline T op(const T &a, const T &b, const T &c) {
return sum::op<T>(mul::op<T>(a, c), b);
}
};
template<> __device__ inline float2 fma_AxCtB::op<float2>(const float2 &a, const float2 &b, const float2 &c) {
#ifdef KITTENS_BLACKWELL
float2 d;
asm volatile("fma.rn.f32x2 %0, %1, %2, %3;" : "=l"(*(uint64_t*)&d) : "l"(*(uint64_t*)&a), "l"(*(uint64_t*)&c), "l"(*(uint64_t*)&b));
return d;
#else
return float2{a.x*c.x+b.x, a.y*c.y+b.y};
#endif
}
} // namespace base_ops
} // namespace kittens
@@ -1,519 +0,0 @@
/**
* @file
* @brief Declarations, manipulations, and wrappers for basic types.
*
* This file is a bunch of utilities for going back and forth between different types.
*
* Many of them are for the compiler, so as to clean up the code. It unfortunately
* seems necessary when we have types we really care about that are less than word width.
*/
#pragma once
#ifdef KITTENS_HOPPER
#include <cuda_fp8.h>
#endif
#include <cuda_bf16.h>
#include <cuda_fp16.h>
#include <string>
#include <bit>
namespace kittens {
/**
* @brief Bfloat16 floating-point type.
*/
using bf16 = __nv_bfloat16;
/**
* @brief Half-precision floating-point type.
*/
using half = __half;
/**
* @brief Packed word of two bfloat16 floating-point values.
*/
using bf16_2 = __nv_bfloat162;
/**
* @brief Packed word of two half-precision floating-point values.
*/
using half_2 = __half2;
#ifdef KITTENS_HOPPER
/**
* @brief float8 floating-point type.
*/
using fp8e4m3 = __nv_fp8_e4m3;
using fp8e5m2 = __nv_fp8_e5m2;
#ifdef KITTENS_BLACKWELL
using fp8e8m0 = __nv_fp8_e8m0;
#endif
/**
* @brief 2-packed float8 floating-point type.
*/
using fp8e4m3_2 = __nv_fp8x2_e4m3;
using fp8e5m2_2 = __nv_fp8x2_e5m2;
#ifdef KITTENS_BLACKWELL
using fp8e8m0_2 = __nv_fp8x2_e8m0;
#endif
/**
* @brief 4-packed float8 floating-point type.
*/
using fp8e4m3_4 = __nv_fp8x4_e4m3;
using fp8e5m2_4 = __nv_fp8x4_e5m2;
#ifdef KITTENS_BLACKWELL
using fp8e8m0_4 = __nv_fp8x4_e8m0;
#endif
#endif
namespace ducks {
/**
* @namespace base_types
*
* @brief A namespace for concepts for basic data types.
*/
namespace base_types {
#ifdef KITTENS_HOPPER
#ifdef KITTENS_BLACKWELL
template<typename T>
concept T2 = std::is_same_v<T, float2> || std::is_same_v<T, bf16_2> || std::is_same_v<T, half_2> || std::is_same_v<T, fp8e4m3_4> || std::is_same_v<T, fp8e5m2_4> || std::is_same_v<T, fp8e8m0_4>; // could add half_2 later if implemented.
template<typename T>
concept T1 = std::is_same_v<T, float> || std::is_same_v<T, bf16 > || std::is_same_v<T, half> || std::is_same_v<T, fp8e4m3> || std::is_same_v<T, fp8e5m2> || std::is_same_v<T, fp8e8m0>; // could add half_2 later if implemented.
#else
template<typename T>
concept T2 = std::is_same_v<T, float2> || std::is_same_v<T, bf16_2> || std::is_same_v<T, half_2> || std::is_same_v<T, fp8e4m3_4> || std::is_same_v<T, fp8e5m2_4>;
template<typename T>
concept T1 = std::is_same_v<T, float> || std::is_same_v<T, bf16 > || std::is_same_v<T, half> || std::is_same_v<T, fp8e4m3> || std::is_same_v<T, fp8e5m2>;
#endif
#else
template<typename T>
concept T2 = std::is_same_v<T, float2> || std::is_same_v<T, bf16_2> || std::is_same_v<T, half_2>;
template<typename T>
concept T1 = std::is_same_v<T, float> || std::is_same_v<T, bf16 > || std::is_same_v<T, half>;
#endif
} // namespace base_types
} // namespace ducks
/**
* @namespace base_types
*
* @brief A namespace for ThunderKittens basic data types.
*/
namespace base_types {
/**
* @brief Provides compile-time constants for different types.
*
* @tparam T The type for which to provide constants.
*/
template<typename T> struct constants {
/**
* @brief Zero
* @return Constexpr zero with type T
*/
static __device__ inline constexpr T zero() { return T{0}; }
/**
* @brief One
* @return Constexpr one with type T
*/
static __device__ inline constexpr T one() { return T{1}; }
/**
* @brief Positive infinity. Particularly useful for initializing before a min op.
* @return Constexpr positive infinity with type T
*/
static __device__ inline constexpr T pos_infty() { return T{INFINITY}; } // I'll find a better way at some point but this appears to work.
/**
* @brief Negative infinity. Particularly useful for initializing before a max op.
* @return Constexpr negative infinity with type T
*/
static __device__ inline constexpr T neg_infty() { return T{-INFINITY}; }
};
template<> struct constants<float2> {
static __device__ inline constexpr float2 zero() { return float2{0.f, 0.f}; }
static __device__ inline constexpr float2 one() { return float2{1.f, 1.f}; }
static __device__ inline constexpr float2 pos_infty() { return float2{constants<float>::pos_infty(), constants<float>::pos_infty()}; }
static __device__ inline constexpr float2 neg_infty() { return float2{constants<float>::neg_infty(), constants<float>::neg_infty()}; }
};
template<> struct constants<bf16> {
static __device__ inline constexpr bf16 zero() { return std::bit_cast<__nv_bfloat16>(uint16_t(0x0000)); } // unfortunately __float2bf16_rn is not constexpr
static __device__ inline constexpr bf16 one() { return std::bit_cast<__nv_bfloat16>(uint16_t(0x3F80)); }
static __device__ inline constexpr bf16 pos_infty() { return std::bit_cast<__nv_bfloat16>(uint16_t(0x7F80)); }
static __device__ inline constexpr bf16 neg_infty() { return std::bit_cast<__nv_bfloat16>(uint16_t(0xFF80)); }
};
template<> struct constants<bf16_2> {
static __device__ inline constexpr bf16_2 zero() { return bf16_2{constants<bf16>::zero(), constants<bf16>::zero()}; }
static __device__ inline constexpr bf16_2 one() { return bf16_2{constants<bf16>::one(), constants<bf16>::one()}; }
static __device__ inline constexpr bf16_2 pos_infty() { return bf16_2{constants<bf16>::pos_infty(), constants<bf16>::pos_infty()}; }
static __device__ inline constexpr bf16_2 neg_infty() { return bf16_2{constants<bf16>::neg_infty(), constants<bf16>::neg_infty()}; }
};
template<> struct constants<half> {
static __device__ inline constexpr half zero() { return std::bit_cast<__half>(uint16_t(0x0000)); }
static __device__ inline constexpr half one() { return std::bit_cast<__half>(uint16_t(0x3C00)); }
static __device__ inline constexpr half pos_infty() { return std::bit_cast<__half>(uint16_t(0x7C00)); }
static __device__ inline constexpr half neg_infty() { return std::bit_cast<__half>(uint16_t(0xFC00)); }
};
template<> struct constants<half_2> {
static __device__ inline constexpr half_2 zero() { return half_2{constants<half>::zero(), constants<half>::zero()}; }
static __device__ inline constexpr half_2 one() { return half_2{constants<half>::one(), constants<half>::one()}; }
static __device__ inline constexpr half_2 pos_infty() { return half_2{constants<half>::pos_infty(), constants<half>::pos_infty()}; }
static __device__ inline constexpr half_2 neg_infty() { return half_2{constants<half>::neg_infty(), constants<half>::neg_infty()}; }
};
#ifdef KITTENS_HOPPER
template<> struct constants<fp8e4m3> {
static __device__ inline constexpr fp8e4m3 zero() { return std::bit_cast<__nv_fp8_e4m3>(uint8_t(0x00)); }
static __device__ inline constexpr fp8e4m3 one() { return std::bit_cast<__nv_fp8_e4m3>(uint8_t(0x38)); }
};
template<> struct constants<fp8e4m3_2> {
static __device__ inline constexpr fp8e4m3_2 zero() { return std::bit_cast<fp8e4m3_2>(uint16_t(0x0000)); }
static __device__ inline constexpr fp8e4m3_2 one() { return std::bit_cast<fp8e4m3_2>(uint16_t(0x3838)); }
};
template<> struct constants<fp8e4m3_4> {
static __device__ inline constexpr fp8e4m3_4 zero() { return std::bit_cast<fp8e4m3_4>(uint32_t(0x00000000)); }
static __device__ inline constexpr fp8e4m3_4 one() { return std::bit_cast<fp8e4m3_4>(uint32_t(0x38383838)); }
};
template<> struct constants<fp8e5m2> {
static __device__ inline constexpr fp8e5m2 zero() { return std::bit_cast<__nv_fp8_e5m2>(uint8_t(0x00)); }
static __device__ inline constexpr fp8e5m2 one() { return std::bit_cast<__nv_fp8_e5m2>(uint8_t(0x3C)); }
};
template<> struct constants<fp8e5m2_2> {
static __device__ inline constexpr fp8e5m2_2 zero() { return std::bit_cast<fp8e5m2_2>(uint16_t(0x0000)); }
static __device__ inline constexpr fp8e5m2_2 one() { return std::bit_cast<fp8e5m2_2>(uint16_t(0x3C3C)); }
};
template<> struct constants<fp8e5m2_4> {
static __device__ inline constexpr fp8e5m2_4 zero() { return std::bit_cast<fp8e5m2_4>(uint32_t(0x00000000)); }
static __device__ inline constexpr fp8e5m2_4 one() { return std::bit_cast<fp8e5m2_4>(uint32_t(0x3C3C3C3C)); }
};
#endif
template<> struct constants<int> {
static __device__ inline constexpr int zero() { return 0; }
static __device__ inline constexpr int one() { return 1; }
};
template<> struct constants<int2> {
static __device__ inline constexpr int2 zero() { return int2{0, 0}; }
static __device__ inline constexpr int2 one() { return int2{1, 1}; }
};
/**
* @brief Provides information about packing of elements for a given type.
*
* @tparam T The type for which to provide packing information.
*/
template<typename T> struct packing {
/**
* @brief The number of elements packed together.
*
* @return constexpr int representing number of elements within the type.
*/
static __device__ inline constexpr int num() { return 1; }
/**
* @brief Packs a single T element twice (replicated) into its packed type.
*
* @param i[in] The element to pack.
* @return The packed type.
*/
static __device__ inline constexpr T pack(const bf16 &i);
};
template<> struct packing<bf16> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = bf16;
using packed_type = bf16_2;
static __device__ inline constexpr bf16_2 pack(const bf16 &i) { return bf16_2{i, i}; }
};
template<> struct packing<bf16_2> {
static __device__ inline constexpr int num() { return 2; }
using unpacked_type = bf16;
using packed_type = bf16_2;
static __device__ inline constexpr bf16_2 pack(const bf16 &i) { return bf16_2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<half> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = half;
using packed_type = half_2;
static __device__ inline constexpr half_2 pack(const half &i) { return half_2{i, i}; }
};
template<> struct packing<half_2> {
static __device__ inline constexpr int num() { return 2; }
using unpacked_type = half;
using packed_type = half_2;
static __device__ inline constexpr half_2 pack(const half &i) { return half_2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<float> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = float;
using packed_type = float2;
static __device__ inline constexpr float2 pack(const float &i) { return float2{i, i}; }
};
template<> struct packing<float2> {
static __device__ inline constexpr int num() { return 2; }
using unpacked_type = float;
using packed_type = float2;
static __device__ inline constexpr float2 pack(const float &i) { return float2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<char> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = char;
using packed_type = char2;
static __device__ inline constexpr char2 pack(const char &i) { return char2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<char2> {
static __device__ inline constexpr int num() { return 2; }
using unpacked_type = char;
using packed_type = char2;
static __device__ inline constexpr char2 pack(const char &i) { return char2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<int> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = int;
using packed_type = int2;
static __device__ inline constexpr int2 pack(const int &i) { return int2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<int2> {
static __device__ inline constexpr int num() { return 2; }
using unpacked_type = int;
using packed_type = int2;
static __device__ inline constexpr int2 pack(const int &i) { return int2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<uint> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = uint;
using packed_type = uint2;
static __device__ inline constexpr uint2 pack(const uint &i) { return uint2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<uint2> {
static __device__ inline constexpr int num() { return 2; }
using unpacked_type = uint;
using packed_type = uint2;
static __device__ inline constexpr uint2 pack(const uint &i) { return uint2{i, i}; } // this replication makes code cleaner later.
};
struct uint64_2 { uint64_t x, y; };
template<> struct packing<uint64_t> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = uint64_t;
using packed_type = uint64_2;
static __device__ inline constexpr uint64_2 pack(const uint64_t &i) { return uint64_2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<uint64_2> {
static __device__ inline constexpr int num() { return 2; }
using unpacked_type = uint64_t;
using packed_type = uint64_2;
static __device__ inline constexpr uint64_2 pack(const uint64_t &i) { return uint64_2{i, i}; } // this replication makes code cleaner later.
};
template<> struct packing<float4> {
static __device__ inline constexpr int num() { return 4; }
};
template<> struct packing<int4> {
static __device__ inline constexpr int num() { return 4; }
};
#ifdef KITTENS_HOPPER
template<> struct packing<fp8e4m3> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = fp8e4m3;
using packed_type = fp8e4m3_4;
};
template<> struct packing<fp8e4m3_4> {
static __device__ inline constexpr int num() { return 4; }
using unpacked_type = fp8e4m3;
using packed_type = fp8e4m3_4;
};
template<> struct packing<fp8e5m2> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = fp8e5m2;
using packed_type = fp8e5m2_4;
};
template<> struct packing<fp8e5m2_4> {
static __device__ inline constexpr int num() { return 4; }
using unpacked_type = fp8e5m2;
using packed_type = fp8e5m2_4;
};
#ifdef KITTENS_BLACKWELL
template<> struct packing<fp8e8m0> {
static __device__ inline constexpr int num() { return 1; }
using unpacked_type = fp8e8m0;
using packed_type = fp8e8m0_4;
};
template<> struct packing<fp8e8m0_4> {
static __device__ inline constexpr int num() { return 4; }
using unpacked_type = fp8e8m0;
using packed_type = fp8e8m0_4;
};
#endif
#endif
/**
* @brief Provides templated functionality to convert between different types.
*
* @tparam T The target type for conversion.
* @tparam U The source type for conversion.
*/
template<typename T, typename U> struct convertor {
/**
* @brief Converts a value of type U to type T.
*
* @param u[in] The value of type U to convert.
* @return T The converted value of type T.
*/
static __host__ __device__ inline T convert(const U & u) {
return (T)u;
}
};
template<> struct convertor<float, bf16> {
static __host__ __device__ inline float convert(const bf16 & u) {
return __bfloat162float(u);
}
};
template<> struct convertor<bf16, float> {
static __host__ __device__ inline bf16 convert(const float & u) {
return __float2bfloat16_rn(u);
}
};
template<> struct convertor<float2, bf16_2> {
static __host__ __device__ inline float2 convert(const bf16_2 & u) {
return __bfloat1622float2(u);
}
};
template<> struct convertor<bf16_2, float2> {
static __host__ __device__ inline bf16_2 convert(const float2 & u) {
return __float22bfloat162_rn(u);
}
};
template<> struct convertor<float, half> {
static __host__ __device__ inline float convert(const half & u) {
return __half2float(u);
}
};
template<> struct convertor<half, float> {
static __host__ __device__ inline half convert(const float & u) {
return __float2half(u);
}
};
template<> struct convertor<float2, half_2> {
static __host__ __device__ inline float2 convert(const half_2 & u) {
return __half22float2(u);
}
};
template<> struct convertor<half_2, float2> {
static __host__ __device__ inline half_2 convert(const float2 & u) {
return __float22half2_rn(u);
}
};
template<> struct convertor<bf16, half> {
static __host__ __device__ inline bf16 convert(const half & u) {
return __float2bfloat16_rn(__half2float(u));
}
};
template<> struct convertor<half, bf16> {
static __host__ __device__ inline half convert(const bf16 & u) {
return __float2half(__bfloat162float(u));
}
};
template<> struct convertor<bf16_2, half_2> {
static __host__ __device__ inline bf16_2 convert(const half_2 & u) {
return __float22bfloat162_rn(__half22float2(u));
}
};
template<> struct convertor<half_2, bf16_2> {
static __host__ __device__ inline half_2 convert(const bf16_2 & u) {
return __float22half2_rn(__bfloat1622float2(u));
}
};
#ifdef KITTENS_HOPPER
// fp8e4m3
template<> struct convertor<fp8e4m3_4, float4> {
static __host__ __device__ inline fp8e4m3_4 convert(const float4& u) {
return __nv_fp8x4_e4m3(u);
}
};
template<> struct convertor<float4, fp8e4m3_4> {
static __host__ __device__ inline float4 convert(const fp8e4m3_4& u) {
__nv_fp8_e4m3 *vals = reinterpret_cast<__nv_fp8_e4m3*>(const_cast<__nv_fp8x4_e4m3*>(&u));
return make_float4(float(vals[0]), float(vals[1]), float(vals[2]), float(vals[3]));
}
};
template<> struct convertor<fp8e4m3_2, float2> {
static __host__ __device__ inline fp8e4m3_2 convert(const float2& u) {
return __nv_fp8x2_e4m3(u);
}
};
template<> struct convertor<float2, fp8e4m3_2> {
static __host__ __device__ inline float2 convert(const fp8e4m3_2& u) {
__nv_fp8_e4m3 *vals = reinterpret_cast<__nv_fp8_e4m3*>(const_cast<__nv_fp8x2_e4m3*>(&u));
return make_float2(float(vals[0]), float(vals[1]));
}
};
template<> struct convertor<fp8e4m3, float> {
static __host__ __device__ inline fp8e4m3 convert(const float & u) {
return __nv_fp8_e4m3(u);
}
};
template<> struct convertor<float, fp8e4m3> {
static __host__ __device__ inline float convert(const fp8e4m3 & u) {
return float(u);
}
};
template<> struct convertor<bf16_2, fp8e4m3_4> {
static __host__ __device__ inline bf16_2 convert(const fp8e4m3_4 & u) {
float4 f4 = convertor<float4, fp8e4m3_4>::convert(u);
float2 f2 = make_float2(f4.x, f4.y);
return __float22bfloat162_rn(f2);
}
};
template<> struct convertor<fp8e4m3_4, bf16_2> {
static __host__ __device__ inline fp8e4m3_4 convert(const bf16_2 & u) {
float2 f2 = __bfloat1622float2(u);
float4 f4 = make_float4(f2.x, f2.y, 0.0f, 0.0f);
return __nv_fp8x4_e4m3(f4);
}
};
// fp8e5m2
template<> struct convertor<fp8e5m2_4, float4> {
static __host__ __device__ inline fp8e5m2_4 convert(const float4& u) {
return __nv_fp8x4_e5m2(u);
}
};
template<> struct convertor<float4, fp8e5m2_4> {
static __host__ __device__ inline float4 convert(const fp8e5m2_4& u) {
__nv_fp8_e5m2 *vals = reinterpret_cast<__nv_fp8_e5m2*>(const_cast<__nv_fp8x4_e5m2*>(&u));
return make_float4(float(vals[0]), float(vals[1]), float(vals[2]), float(vals[3]));
}
};
template<> struct convertor<fp8e5m2_2, float2> {
static __host__ __device__ inline fp8e5m2_2 convert(const float2& u) {
return __nv_fp8x2_e5m2(u);
}
};
template<> struct convertor<float2, fp8e5m2_2> {
static __host__ __device__ inline float2 convert(const fp8e5m2_2& u) {
__nv_fp8_e5m2 *vals = reinterpret_cast<__nv_fp8_e5m2*>(const_cast<__nv_fp8x2_e5m2*>(&u));
return make_float2(float(vals[0]), float(vals[1]));
}
};
template<> struct convertor<fp8e5m2, float> {
static __host__ __device__ inline fp8e5m2 convert(const float & u) {
return __nv_fp8_e5m2(u);
}
};
template<> struct convertor<float, fp8e5m2> {
static __host__ __device__ inline float convert(const fp8e5m2 & u) {
return float(u);
}
};
template<> struct convertor<bf16_2, fp8e5m2_4> {
static __host__ __device__ inline bf16_2 convert(const fp8e5m2_4 & u) {
float4 f4 = convertor<float4, fp8e5m2_4>::convert(u);
float2 f2 = make_float2(f4.x, f4.y);
return __float22bfloat162_rn(f2);
}
};
template<> struct convertor<fp8e5m2_4, bf16_2> {
static __host__ __device__ inline fp8e5m2_4 convert(const bf16_2 & u) {
float2 f2 = __bfloat1622float2(u);
float4 f4 = make_float4(f2.x, f2.y, 0.0f, 0.0f);
return __nv_fp8x4_e5m2(f4);
}
};
#endif
}
}
@@ -1,11 +0,0 @@
/**
* @file
* @brief A collection of common resources on which ThunderKittens depends.
*/
#pragma once
#include "util.cuh"
#include "base_types.cuh"
#include "base_ops.cuh"
@@ -1,56 +0,0 @@
#pragma once
// Reset
#define TK_RESET "\033[0m"
// Foreground colors
#define TK_FG_BLACK "\033[30m"
#define TK_FG_RED "\033[31m"
#define TK_FG_GREEN "\033[32m"
#define TK_FG_YELLOW "\033[33m"
#define TK_FG_BLUE "\033[34m"
#define TK_FG_MAGENTA "\033[35m"
#define TK_FG_CYAN "\033[36m"
#define TK_FG_WHITE "\033[37m"
// Background colors
#define TK_BG_BLACK "\033[40m"
#define TK_BG_RED "\033[41m"
#define TK_BG_GREEN "\033[42m"
#define TK_BG_YELLOW "\033[43m"
#define TK_BG_BLUE "\033[44m"
#define TK_BG_MAGENTA "\033[45m"
#define TK_BG_CYAN "\033[46m"
#define TK_BG_WHITE "\033[47m"
// Bright foreground colors
#define TK_FG_BRIGHT_BLACK "\033[90m"
#define TK_FG_BRIGHT_RED "\033[91m"
#define TK_FG_BRIGHT_GREEN "\033[92m"
#define TK_FG_BRIGHT_YELLOW "\033[93m"
#define TK_FG_BRIGHT_BLUE "\033[94m"
#define TK_FG_BRIGHT_MAGENTA "\033[95m"
#define TK_FG_BRIGHT_CYAN "\033[96m"
#define TK_FG_BRIGHT_WHITE "\033[97m"
// Bright background colors
#define TK_BG_BRIGHT_BLACK "\033[100m"
#define TK_BG_BRIGHT_RED "\033[101m"
#define TK_BG_BRIGHT_GREEN "\033[102m"
#define TK_BG_BRIGHT_YELLOW "\033[103m"
#define TK_BG_BRIGHT_BLUE "\033[104m"
#define TK_BG_BRIGHT_MAGENTA "\033[105m"
#define TK_BG_BRIGHT_CYAN "\033[106m"
#define TK_BG_BRIGHT_WHITE "\033[107m"
// Text styles
#define TK_BOLD "\033[1m"
#define TK_DIM "\033[2m"
#define TK_ITALIC "\033[3m"
#define TK_UNDERLINE "\033[4m"
#define TK_BLINK "\033[5m"
#define TK_REVERSE "\033[7m"
#define TK_HIDDEN "\033[8m"
// Macro to combine styles
#define TK_STYLE(...) "\033[" #__VA_ARGS__ "m"
-314
View File
@@ -1,314 +0,0 @@
/**
* @file
* @brief General utilities for ThunderKittens.
*/
#pragma once
#include <stdint.h>
#include <type_traits>
#include <concepts>
#include <memory>
// CUDA driver API
#define CUCHECK(cmd) do { \
CUresult err = cmd; \
if (err != CUDA_SUCCESS) { \
const char *errStr; \
cuGetErrorString(err, &errStr); \
fprintf(stderr, "Failed: CUDA error %s:%d '%s'\n", \
__FILE__, __LINE__, errStr); \
exit(EXIT_FAILURE); \
} \
} while(0)
// CUDA runtime API
#define CUDACHECK(cmd) do { \
cudaError_t err = cmd; \
if (err != cudaSuccess) { \
fprintf(stderr, "Failed: CUDA error %s:%d '%s'\n", \
__FILE__, __LINE__, cudaGetErrorString(err)); \
exit(EXIT_FAILURE); \
} \
} while(0)
/**
* @namespace kittens
*
* @brief The main namespace of ThunderKittens.
*/
namespace kittens {
/* ---------- GENERAL CONSTANTS FOR KITTENS ---------- */
/**
* @brief Tile dimension constant.
*/
template<typename T> constexpr int TILE_COL_DIM = sizeof(T) == 1 ? 32 : 16;
template<typename T> constexpr int TILE_ROW_DIM = 16;
/**
* @brief Tile num elements constant calculated as TILE_DIM squared.
*/
template<typename T> constexpr int TILE_ELEMENTS{TILE_COL_DIM<T>*TILE_ROW_DIM<T>};
/**
* @brief Constant representing number of threads in a warp.
*/
constexpr int WARP_THREADS{32};
/**
* @brief Constant representing number of threads in a warpgroup of four warps.
*/
constexpr int WARPGROUP_THREADS{128};
/**
* @brief Constant representing number of warps in a warpgroup of four warps.
*/
constexpr int WARPGROUP_WARPS{4};
/**
* @brief Get the warp ID of the current thread.
* @return The warp ID.
*/
__device__ static __forceinline__ int warpid() {
// uint32_t wid;
// asm volatile("mov.u32 %0, %warpid;" : "=r"(wid));
// return wid;
return threadIdx.x >> 5;
}
/**
* @brief Get the warpgroup ID of the current thread.
* @return The warpgroup ID.
*/
__device__ static __forceinline__ int warpgroupid() { return warpid() >> 2; }
/**
* @brief Get the lane ID of the current thread within its warp.
* @return The lane ID.
*/
__device__ static __forceinline__ int laneid() {
// uint32_t lid;
// asm volatile("mov.u32 %0, %laneid;" : "=r"(lid));
// return lid;
return threadIdx.x & 31;
}
#if defined(KITTENS_HOPPER)
constexpr int MAX_SHARED_MEMORY = 227000;
#elif defined(KITTENS_A100)
constexpr int MAX_SHARED_MEMORY = 164000;
#elif defined(KITTENS_4090)
constexpr int MAX_SHARED_MEMORY = 100000;
#endif
struct transpose {
static constexpr int N = 0; // not transposed
static constexpr int T = 1; // transposed
};
struct axis {
static constexpr int ROW = 0; // row axis of a tile
static constexpr int COL = 1; // column axis of a tile
};
/* ---------- TYPE HELPERS ---------- */
/**
* @namespace ducks
*
* @brief ThunderKittens' namespace for template metaprogramming..
*
* This includes primarily dummy types and concept wrappers, along
* with a few additional utilities.
*/
namespace ducks {
/**
* @brief A type representing an empty default for a template.
*/
struct default_type {};
// This macro can't be done as a template, so it doesn't really have a location in kittens.
#define typeof(A) typename std::remove_const<typename std::remove_reference<decltype(A)>::type>::type
}
/* ---------- SHUFFLE UTILS ---------- */
/**
* @brief Mask constant for all active threads in a warp.
*/
static constexpr uint32_t MASK_ALL = 0xFFFFFFFF;
/**
* @brief Perform a shuffle down operation on a packed type synchronously across a warp.
* @tparam T The type of the value to be shuffled.
* @param mask[in] The mask of active threads.
* @param f[in] The value to be shuffled.
* @param delta[in] The number of positions to shuffle down.
* @return The result of the shuffle operation.
*/
template<typename T>
__device__ static inline T packed_shfl_down_sync(uint32_t mask, const T &f, int delta) {
return __shfl_down_sync(mask, f, delta);
}
template<>
__device__ inline float2 packed_shfl_down_sync<float2>(uint32_t mask, const float2 &f, int delta) {
float2 r;
r.x = __shfl_down_sync(mask, f.x, delta);
r.y = __shfl_down_sync(mask, f.y, delta);
return r;
}
/**
* @brief Perform a packed shuffle operation synchronously across a warp.
* @tparam T The type of the value to be shuffled.
* @param mask[in] The mask of active threads.
* @param f[in] The value to be shuffled.
* @param src[in] The source lane from which to shuffle.
* @return The result of the shuffle operation.
*/
template<typename T>
__device__ static inline T packed_shfl_sync(uint32_t mask, const T &f, int src) {
return __shfl_sync(mask, f, src);
}
template<>
__device__ inline float2 packed_shfl_sync<float2>(uint32_t mask, const float2 &f, int src) {
float2 r;
r.x = __shfl_sync(mask, f.x, src);
r.y = __shfl_sync(mask, f.y, src);
return r;
}
/* ---------- SHARED MEMORY UTILS ---------- */
// namespace ducks {
// namespace sb {
// struct identifier {};
// }
// }
// template<typename Args...>
// struct sb {
// using identifier = ducks::sb::identifier;
// Args... args;
// };
// namespace ducks {
// namespace sb {
// template<typename T> concept all = requires {
// typename T::identifier;
// } && std::is_same_v<T::identifier, identifier>;
// }
// }
// Joyously stolen from https://github.com/NVIDIA/cutlass/blob/5c447dd84f8ae0e1d48ff9a2eae26ce8c4958101/include/cute/container/alignment.hpp#L51
#if defined(__CUDACC__)
#define KITTENS_ALIGN_AS(n) __align__(n)
#else
#define KITTENS_ALIGN_AS(n) alignas(n)
#endif
#ifdef KITTENS_HOPPER
#define KITTENS_DEFAULT_ALIGN KITTENS_ALIGN_AS(128)
#else
#define KITTENS_DEFAULT_ALIGN KITTENS_ALIGN_AS(16)
#endif
/**
* @brief Dummy structure for alignment purposes. Needed for WGMMA and TMA calls.
*/
struct KITTENS_DEFAULT_ALIGN alignment_dummy { int dummy; };
/**
* @brief Very simple allocator for dynamic shared memory. Advances pointer and tracks alignments.
* @tparam default_alignment The default alignment this allocator will enforce. If <=0 (default -1) it will not align.
*/
#ifdef KITTENS_HOPPER
template<int default_alignment=1024>
#else
template<int default_alignment=16>
#endif
struct shared_allocator {
int *ptr;
private:
// Recursive template to generate N-dimensional array type
template<typename A, size_t... dims>
struct variadic_array;
template<typename A, size_t first_dim, size_t... rest_dims>
struct variadic_array<A, first_dim, rest_dims...> {
using type = typename variadic_array<A, rest_dims...>::type[first_dim];
};
template<typename A>
struct variadic_array<A> {
using type = A;
};
template<typename A, size_t... dims>
using variadic_array_t = typename variadic_array<A, dims...>::type;
template<int alignment>
__device__ inline void align_ptr() {
if constexpr (alignment > 0) {
uint64_t p = reinterpret_cast<uint64_t>(ptr);
if(p % alignment != 0) {
ptr = (int*)(p + (alignment-(p%alignment)));
}
}
}
public:
/**
* @brief Construct a new shared allocator using a pointer to extern shared memory.
* @param[in] _ptr Pointer to the start of the extern shared memory.
*/
__device__ shared_allocator(int *_ptr): ptr(_ptr) {}
/**
* @brief Allocate shared memory for a single instance or N-dimensional array of type A.
* @tparam A The type of the object to allocate.
* @tparam dims... A list of dimensions for the N-dimensional array.
* @return Reference to the allocated object.
*/
template<typename A, size_t... dims>
__device__ inline variadic_array_t<A, dims...>& allocate() {
// static_assert(sizeof(A) % default_alignment == 0, "Type is not aligned properly for array allocation");
align_ptr<default_alignment>();
using at = variadic_array_t<A, dims...>;
at*p = reinterpret_cast<at*>(ptr);
ptr += sizeof(at)/sizeof(int);
return *p;
}
/**
* @brief Allocate shared memory for a single instance or N-dimensional array of type A.
* @tparam alignment An alignment to enforce for this particular object.
* @tparam A The type of the object to allocate.
* @tparam dims... A list of dimensions for the N-dimensional array.
* @return Reference to the allocated object.
*/
template<int alignment, typename A, size_t... dims>
__device__ inline variadic_array_t<A, dims...>& allocate() {
// static_assert(sizeof(A) % alignment == 0, "Type is not aligned properly for array allocation");
align_ptr<alignment>();
using at = variadic_array_t<A, dims...>;
at*p = reinterpret_cast<at*>(ptr);
ptr += sizeof(at)/sizeof(int);
return *p;
}
};
#if (defined(KITTENS_HOPPER) || defined(KITTENS_BLACKWELL))
/**
* @brief A wrapper for an allocator that enforces sufficient alignment to be used for TMA loads and stores.
*/
using tma_allocator = shared_allocator<1024>;
using tma_swizzle_allocator = tma_allocator; // swizzled TMA modes require up to 1024 byte alignments :/
/* Get CTA ID within a cluster */
__device__ static inline int3 clusterIdx() {
int3 cluster_idx;
asm volatile("mov.u32 %0, %clusterid.x;\n" : "=r"(cluster_idx.x));
asm volatile("mov.u32 %0, %clusterid.y;\n" : "=r"(cluster_idx.y));
asm volatile("mov.u32 %0, %clusterid.z;\n" : "=r"(cluster_idx.z));
return cluster_idx;
}
__device__ static inline int cluster_ctarank() {
uint32_t ctarank;
asm volatile("mov.u32 %0, %cluster_ctarank;\n" : "=r"(ctarank));
return ctarank;
}
#endif
} // namespace kittens
-12
View File
@@ -1,12 +0,0 @@
/**
* @file
* @brief The master header file of ThunderKittens. This file includes everything you need!
*/
#pragma once
#include "common/common.cuh"
#include "types/types.cuh"
#include "ops/ops.cuh"
#include "pyutils/util.cuh"
// #include "pyutils/pyutils.cuh" // for simple binding without including torch
@@ -1,51 +0,0 @@
/**
* @file
* @brief An aggregate header of all device (multi-GPU) operations defined by ThunderKittens
*/
#pragma once
#include "../../types/types.cuh"
namespace kittens {
template<int _NUM_DEVICES>
struct device {
static_assert(_NUM_DEVICES >= 0 && _NUM_DEVICES <= 72, "Invalid number of devices");
static constexpr int NUM_DEVICES = _NUM_DEVICES;
#ifdef KITTENS_HOPPER
using barrier_t = pgl<gl<int, 1, 1, 1, -1>, NUM_DEVICES, true>;
/**
* @brief Multi-GPU synchronization barrier for coordinated kernel exit
*
* Performs a synchronization across all devices to ensure all GPUs complete
* their work before any kernel exits. Does not synchronize intra-node threads
* or threadblocks.
*
* @param barrier Pre-allocated barrier structure, must be initialized to 0
* @param dev_idx Current device index (0 to NUM_DEVICES - 1)
* @param id Synchronization point identifier (default: 0). 0 is fine for most cases
*
*/
__device__ static inline void sync_on_exit(const barrier_t &barrier, const int dev_idx, const int id = 0) {
if (blockIdx.x == 0 && blockIdx.y == 0 && blockIdx.z == 0 &&
threadIdx.x == 0 && threadIdx.y == 0 && threadIdx.z == 0) {
cuda::atomic_ref<int, cuda::thread_scope_system> barrier_uc(barrier[dev_idx][{id}]);
// Inter-note check-in
multimem<int>::red<reduce_op::ADD>(barrier.mc_ptr_at({id}), 1);
asm volatile ("{fence.proxy.alias;}" ::: "memory");
while (barrier_uc.load(cuda::memory_order_acquire) < NUM_DEVICES);
barrier_uc.fetch_sub(NUM_DEVICES, cuda::memory_order_release);
}
}
#endif
};
} // namespace kittens
@@ -1,96 +0,0 @@
/**
* @file
* @brief An aggregate header of all group (multi-warp) operations defined by ThunderKittens
*/
#pragma once
#include <cuda/pipeline>
#include "../../common/common.cuh"
#include "../../types/types.cuh"
#include "../thread/thread.cuh" // several group memory ops rely on underlying warp-scope ops
#define KITTENS_CHECK_WARP static_assert(GROUP_WARPS==1, "Warp (GROUP_WARPS=1) function called from a non-warp group.");
// A "warpgroup" is a special group of 4 consecutive warps defined by NVIDIA for certain SM_90+ operations.
#define KITTENS_CHECK_WARPGROUP static_assert(GROUP_WARPS==4, "Warpgroup (GROUP_WARPS=4) function called from a non-warpgroup group.");
// WGMMA relies on some template structures that cannot be specialized within the group struct, so we declare them in advance.
#ifdef KITTENS_HOPPER
#include "mma/warpgroup/base/base.cuh"
#endif
namespace kittens {
/*
This is meant to be used with a `using group_N = kittens::group<NUM_WORKERS>;` at the start of every kernel.
*/
template<int _GROUP_WARPS>
struct group {
static constexpr int GROUP_WARPS = _GROUP_WARPS; // This alias produces nice parallelism.
static constexpr int GROUP_THREADS = GROUP_WARPS * kittens::WARP_THREADS; // This alias produces nice parallelism.
__device__ static inline int laneid() { return threadIdx.x % GROUP_THREADS; }
__device__ static inline int warpid() { return laneid() / kittens::WARP_THREADS; }
__device__ static inline int groupid() { return threadIdx.x / GROUP_THREADS; }
__device__ static inline void sync(int id) {
asm volatile("bar.sync %0, %1;\n" :: "r"(id), "n"(GROUP_THREADS));
}
template<uint32_t MASK=0xFFFFFFFF> __device__ static inline void sync() {
static_assert(GROUP_WARPS==1, "barrier-less sync() can only be called by a single warp!");
asm volatile("bar.warp.sync %0;\n" :: "n"(MASK));
}
__device__ static inline void arrive(int id) {
asm volatile("bar.arrive %0, %1;\n" :: "r"(id), "n"(GROUP_THREADS));
}
#include "memory/memory.cuh"
#include "shared/shared.cuh"
#include "register/register.cuh"
#ifdef KITTENS_HOPPER
#include "mma/mma.cuh"
template<int n_reg> __device__ static inline void increase_registers() {
static_assert(n_reg % 8 == 0, "n_reg must be a multiple of 8");
asm volatile("setmaxnreg.inc.sync.aligned.u32 %0;\n" :: "n"(n_reg));
}
template<int n_reg> __device__ static inline void decrease_registers() {
static_assert(n_reg % 8 == 0, "n_reg must be a multiple of 8");
asm volatile("setmaxnreg.dec.sync.aligned.u32 %0;\n" :: "n"(n_reg));
}
__device__ static inline void producer_registers() { decrease_registers<24>(); }
template<int NCWG> __device__ static inline void consumer_registers() { increase_registers<480/NCWG - 8*(NCWG>3) - 224*(NCWG==1)>(); }
#endif
};
namespace everyone {
// Block-level synchronization
__device__ static inline void sync(int id) {
asm volatile("bar.sync %0;\n" :: "r"(id));
}
// Cluster-level synchronization functions
namespace tma {
namespace cluster {
__device__ static inline void arrive_aligned() { // All threads in the cluster must call this
asm volatile ("barrier.cluster.arrive.release.aligned;\n");
}
__device__ static inline void wait_aligned() {
asm volatile ("barrier.cluster.wait.acquire.aligned;\n");
}
__device__ static inline void sync() {
arrive_aligned();
wait_aligned();
}
}
}
};
using warp = group<1>; // scope used by most pre-Hopper GPUs, and also for most register operations.
using warpgroup = group<4>; // special scope commonly used by Hopper and later.
}
@@ -1,21 +0,0 @@
/**
* @file
* @brief An aggregate header of colaborative group memory movement operations
*/
#include "util/util.cuh"
#include "tile/tile.cuh"
#include "vec/vec.cuh"
#ifdef KITTENS_HOPPER
struct tma {
#include "util/tma.cuh"
#include "tile/tma.cuh"
#include "vec/tma.cuh"
struct cluster {
#include "util/tma_cluster.cuh"
#include "tile/tma_cluster.cuh"
#include "vec/tma_cluster.cuh"
};
};
#endif
@@ -1,42 +0,0 @@
/**
* @file
* @brief Functions for a group to collaboratively transfer data directly between global memory and registers and back.
*/
/**
* @brief Collaboratively loads data from a source array into register tiles.
*
* @tparam RT The register tile type.
* @tparam U The data type of the source array.
* @param dst[out] The destination tile to load data into.
* @param src[in] The source array to load data from.
* @param row_stride[in] The stride in elements between rows in the source array.
*/
template<int axis, ducks::crt::all CRT, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<crt<typename CRT::T, GROUP_WARPS*CRT::rows, CRT::cols, typename CRT::layout>>>
__device__ inline static void load(CRT &dst, const CGL &src, const COORD &idx) {
load<axis, CRT::component, CGL::component, COORD>(dst.real, src.real, idx);
load<axis, CRT::component, CGL::component, COORD>(dst.imag, src.imag, idx);
}
template<ducks::crt::all CRT, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<crt<typename CRT::T, GROUP_WARPS*CRT::rows, CRT::cols, typename CRT::layout>>>
__device__ inline static void load(CRT &dst, const CGL &src, const COORD &idx) {
load<2, CRT, CGL>(dst, src, idx);
}
/**
* @brief Collaboratively stores data from register tiles to a destination array in global memory.
*
* @tparam RT The register tile type.
* @tparam U The data type of the destination array.
* @param[out] dst The destination array in global memory to store data into.
* @param[in] src The source register tile to store data from.
* @param row_stride[in] The stride in elements between rows in the destination array.
*/
template<int axis, ducks::crt::all CRT, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<crt<typename CRT::T, GROUP_WARPS*CRT::rows, CRT::cols, typename CRT::layout>>>
__device__ inline static void store(CGL &dst, const CRT &src, const COORD &idx) {
store<axis, typename CRT::component, typename CGL::component>(dst.real, src.real, idx);
store<axis, typename CRT::component, typename CGL::component>(dst.imag, src.imag, idx);
}
template<ducks::crt::all CRT, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<crt<typename CRT::T, GROUP_WARPS*CRT::rows, CRT::cols, typename CRT::layout>>>
__device__ inline static void store(CGL &dst, const CRT &src, const COORD &idx) {
store<2, CRT, CGL>(dst, src, idx);
}
@@ -1,37 +0,0 @@
/**
* @file
* @brief Group (collaborative warp) ops for loading shared tiles from and storing to global memory.
*/
template<int axis, bool assume_aligned, ducks::cst::all CST, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<CST>>
__device__ static inline void load(CST &dst, const CGL &src, const COORD &idx) {
load<axis, assume_aligned, typename CST::component, typename CGL::component, COORD>(dst.real, src.real, idx);
load<axis, assume_aligned, typename CST::component, typename CGL::component, COORD>(dst.imag, src.imag, idx);
}
template<ducks::cst::all CST, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<CST>>
__device__ static inline void load(CST &dst, const CGL &src, const COORD &idx) {
load<2, false, typename CST::component, typename CGL::component, COORD>(dst.real, src.real, idx);
load<2, false, typename CST::component, typename CGL::component, COORD>(dst.imag, src.imag, idx);
}
template<int axis, bool assume_aligned, ducks::cst::all CST, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<CST>>
__device__ static inline void store(CGL &dst, const CST &src, const COORD &idx) {
store<axis, assume_aligned, typename CST::component, typename CGL::component, COORD>(dst.real, src.real, idx);
store<axis, assume_aligned, typename CST::component, typename CGL::component, COORD>(dst.imag, src.imag, idx);
}
template<ducks::cst::all CST, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<CST>>
__device__ static inline void store(CGL &dst, const CST &src, const COORD &idx) {
store<2, false, typename CST::component, typename CGL::component, COORD>(dst.real, src.real, idx);
store<2, false, typename CST::component, typename CGL::component, COORD>(dst.imag, src.imag, idx);
}
template<int axis, bool assume_aligned, ducks::cst::all CST, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<CST>>
__device__ static inline void load_async(CST &dst, const CGL &src, const COORD &idx) {
load_async<axis, assume_aligned, typename CST::component, typename CGL::component, COORD>(dst.real, src.real, idx);
load_async<axis, assume_aligned, typename CST::component, typename CGL::component, COORD>(dst.imag, src.imag, idx);
}
template<ducks::cst::all CST, ducks::cgl::all CGL, ducks::coord::tile COORD=coord<CST>>
__device__ static inline void load_async(CST &dst, const CGL &src, const COORD &idx) {
load_async<2, false, typename CST::component, typename CGL::component, COORD>(dst.real, src.real, idx);
load_async<2, false, typename CST::component, typename CGL::component, COORD>(dst.imag, src.imag, idx);
}
@@ -1,34 +0,0 @@
/**
* @file
* @brief Functions for a warpgroup to collaboratively transfer data directly between shared memory and registers and back.
*/
/**
* @brief Collaboratively load data from a shared tile into register tiles split across a warpgroup.
*
* @tparam RT The register tile type
* @tparam ST The shared tile type
* @param dst[out] The destination register tile.
* @param src[in] The source shared tile.
*/
template<ducks::crt::all RT, ducks::cst::all ST>
__device__ inline static void load(RT &dst, const ST &src) {
load(dst.real, src.real);
load(dst.imag, src.imag);
}
/**
* @brief Collaboratively store data into a shared tile from register tiles split across a warpgroup.
*
* @tparam RT The register tile type
* @tparam ST The shared tile type
* @param dst[out] The destination shared tile.
* @param src[in] The source register tile.
*/
template<ducks::cst::all ST, ducks::crt::all RT>
__device__ inline static void store(ST &dst, const RT &src) {
store(dst.real, src.real);
store(dst.imag, src.imag);
}
@@ -1,207 +0,0 @@
/**
* @file
* @brief Functions for a group to collaboratively transfer data directly between global memory and registers and back.
*/
/**
* @brief Collaboratively loads data from a source array into row-major layout tiles.
*
* @tparam RT The row-major layout tile type.
* @tparam U The data type of the source array.
* @param dst[out] The destination tile to load data into.
* @param src[in] The source array to load data from.
* @param row_stride[in] The stride in elements between rows in the source array.
*/
template<int axis, ducks::rt::row_layout RT, ducks::gl::all GL, ducks::coord::tile COORD=coord<rt<typename RT::T, GROUP_WARPS*RT::rows, RT::cols, typename RT::layout>>>
__device__ inline static void load(RT &dst, const GL &src, const COORD &idx) {
using T2 = RT::dtype;
using U = typename GL::dtype;
#ifdef KITTENS_HOPPER
static_assert(!std::is_same_v<T2, fp8e4m3_4> && !std::is_same_v<T2, fp8e5m2_4>, "Unsupported type for load/store");
#endif
U *src_ptr = (U*)&src[(idx.template unit_coord<axis, 3>())];
const int row_stride = src.template stride<axis>();
using U2 = base_types::packing<U>::packed_type;
int warp_laneid = threadIdx.x % WARP_THREADS;
int local_warpid;
if constexpr(GROUP_WARPS % 4 == 0) local_warpid = (warpid()/4+(warpid()%4)*(GROUP_WARPS/4));
else local_warpid = warpid();
const int row_offset = dst.rows*local_warpid;
#pragma unroll
for(int i = 0; i < dst.height; i++) {
int row = row_offset + i*dst.tile_size_row + (warp_laneid / 4);
#pragma unroll
for(int j = 0; j < dst.width; j++) {
int col = j*dst.tile_size_col + 2*(warp_laneid % 4);
dst.tiles[i][j].data[0] = base_types::convertor<T2, U2>::convert(*(U2*)(&src_ptr[(row+0)*row_stride + (col+0)]));
dst.tiles[i][j].data[2] = base_types::convertor<T2, U2>::convert(*(U2*)(&src_ptr[(row+0)*row_stride + (col+8)]));
}
#pragma unroll
for(int j = 0; j < dst.width; j++) {
int col = j*dst.tile_size_col + 2*(warp_laneid % 4);
dst.tiles[i][j].data[1] = base_types::convertor<T2, U2>::convert(*(U2*)(&src_ptr[(row+8)*row_stride + (col+0)]));
dst.tiles[i][j].data[3] = base_types::convertor<T2, U2>::convert(*(U2*)(&src_ptr[(row+8)*row_stride + (col+8)]));
}
}
}
/**
* @brief Collaboratively loads data from a source array into column-major layout tiles.
*
* @tparam RT The column-major layout tile type.
* @tparam U The data type of the source array.
* @param dst[out] The destination tile to load data into.
* @param src[in] The source array to load data from.
* @param row_stride[in] The stride in elements between rows in the source array.
*/
template<int axis, ducks::rt::col_layout RT, ducks::gl::all GL, ducks::coord::tile COORD=coord<rt<typename RT::T, GROUP_WARPS*RT::rows, RT::cols, typename RT::layout>>>
__device__ inline static void load(RT &dst, const GL &src, const COORD &idx) {
using T = typename RT::T;
using U = typename GL::dtype;
#ifdef KITTENS_HOPPER
static_assert(!std::is_same_v<T, fp8e4m3> && !std::is_same_v<T, fp8e5m2>, "Unsupported type for load/store");
#endif
U *src_ptr = (U*)&src[(idx.template unit_coord<axis, 3>())];
const int row_stride = src.template stride<axis>();
int warp_laneid = threadIdx.x % WARP_THREADS;
int local_warpid;
if constexpr(GROUP_WARPS % 4 == 0) local_warpid = (warpid()/4+(warpid()%4)*(GROUP_WARPS/4));
else local_warpid = warpid();
const int row_offset = dst.rows*local_warpid;
#pragma unroll
for(int i = 0; i < dst.height; i++) {
int row = row_offset + i*dst.tile_size_row + 2*(warp_laneid % 4);
#pragma unroll
for(int j = 0; j < dst.width; j++) {
int col = j*dst.tile_size_col + (warp_laneid / 4);
dst.tiles[i][j].data[0].x = base_types::convertor<T, U>::convert(src_ptr[(row+0)*row_stride + (col+0)]);
dst.tiles[i][j].data[1].x = base_types::convertor<T, U>::convert(src_ptr[(row+0)*row_stride + (col+8)]);
}
#pragma unroll
for(int j = 0; j < dst.width; j++) {
int col = j*dst.tile_size_col + (warp_laneid / 4);
dst.tiles[i][j].data[0].y = base_types::convertor<T, U>::convert(src_ptr[(row+1)*row_stride + (col+0)]);
dst.tiles[i][j].data[1].y = base_types::convertor<T, U>::convert(src_ptr[(row+1)*row_stride + (col+8)]);
}
#pragma unroll
for(int j = 0; j < dst.width; j++) {
int col = j*dst.tile_size_col + (warp_laneid / 4);
dst.tiles[i][j].data[2].x = base_types::convertor<T, U>::convert(src_ptr[(row+8)*row_stride + (col+0)]);
dst.tiles[i][j].data[3].x = base_types::convertor<T, U>::convert(src_ptr[(row+8)*row_stride + (col+8)]);
}
#pragma unroll
for(int j = 0; j < dst.width; j++) {
int col = j*dst.tile_size_col + (warp_laneid / 4);
dst.tiles[i][j].data[2].y = base_types::convertor<T, U>::convert(src_ptr[(row+9)*row_stride + (col+0)]);
dst.tiles[i][j].data[3].y = base_types::convertor<T, U>::convert(src_ptr[(row+9)*row_stride + (col+8)]);
}
}
}
template<ducks::rt::all RT, ducks::gl::all GL, ducks::coord::tile COORD=coord<rt<typename RT::T, GROUP_WARPS*RT::rows, RT::cols, typename RT::layout>>>
__device__ inline static void load(RT &dst, const GL &src, const COORD &idx) {
load<2>(dst, src, idx);
}
/**
* @brief Collaboratively stores data from register tiles to a destination array in global memory with a row-major layout.
*
* @tparam RT The register tile type with a row-major layout.
* @tparam U The data type of the destination array.
* @param[out] dst The destination array in global memory to store data into.
* @param[in] src The source register tile to store data from.
* @param row_stride[in] The stride in elements between rows in the destination array.
*/
template<int axis, ducks::rt::row_layout RT, ducks::gl::all GL, ducks::coord::tile COORD=coord<rt<typename RT::T, GROUP_WARPS*RT::rows, RT::cols, typename RT::layout>>>
__device__ inline static void store(const GL &dst, const RT &src, const COORD &idx) {
using T2 = RT::dtype;
using U = typename GL::dtype;
#ifdef KITTENS_HOPPER
static_assert(!std::is_same_v<T2, fp8e4m3_4> && !std::is_same_v<T2, fp8e5m2_4>, "Unsupported type for load/store");
#endif
U *dst_ptr = (U*)&dst[(idx.template unit_coord<axis, 3>())];
const int row_stride = dst.template stride<axis>();
using U2 = base_types::packing<U>::packed_type;
int warp_laneid = threadIdx.x % WARP_THREADS;
int local_warpid;
if constexpr(GROUP_WARPS % 4 == 0) local_warpid = (warpid()/4+(warpid()%4)*(GROUP_WARPS/4));
else local_warpid = warpid();
const int row_offset = src.rows*local_warpid;
#pragma unroll
for(int i = 0; i < src.height; i++) {
int row = row_offset + i*src.tile_size_row + (warp_laneid / 4);
#pragma unroll
for(int j = 0; j < src.width; j++) {
int col = j*src.tile_size_col + 2*(warp_laneid % 4);
*(U2*)(&dst_ptr[(row+0)*row_stride + (col+0)]) = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[0]);
*(U2*)(&dst_ptr[(row+0)*row_stride + (col+8)]) = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[2]);
}
#pragma unroll
for(int j = 0; j < src.width; j++) {
int col = j*src.tile_size_col + 2*(warp_laneid % 4);
*(U2*)(&dst_ptr[(row+8)*row_stride + (col+0)]) = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[1]);
*(U2*)(&dst_ptr[(row+8)*row_stride + (col+8)]) = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[3]);
}
}
}
/**
* @brief Collaboratively stores data from register tiles to a destination array in global memory with a column-major layout.
*
* @tparam RT The register tile type with a column-major layout.
* @tparam U The data type of the destination array.
* @param[out] dst The destination array in global memory to store data into.
* @param[in] src The source register tile to store data from.
* @param row_stride[in] The stride in elements between rows in the destination array.
*/
template<int axis, ducks::rt::col_layout RT, ducks::gl::all GL, ducks::coord::tile COORD=coord<rt<typename RT::T, GROUP_WARPS*RT::rows, RT::cols, typename RT::layout>>>
__device__ inline static void store(const GL &dst, const RT &src, const COORD &idx) {
using T = base_types::packing<typename RT::dtype>::unpacked_type;
using U = typename GL::dtype;
#ifdef KITTENS_HOPPER
static_assert(!std::is_same_v<T, fp8e4m3_4> && !std::is_same_v<T, fp8e5m2_4>, "Unsupported type for load/store");
#endif
U *dst_ptr = (U*)&dst[(idx.template unit_coord<axis, 3>())];
const int row_stride = dst.template stride<axis>();
int warp_laneid = threadIdx.x % WARP_THREADS;
int local_warpid;
if constexpr(GROUP_WARPS % 4 == 0) local_warpid = (warpid()/4+(warpid()%4)*(GROUP_WARPS/4));
else local_warpid = warpid();
const int row_offset = src.rows*local_warpid;
#pragma unroll
for(int i = 0; i < src.height; i++) {
int row = row_offset + i*src.tile_size_row + 2*(warp_laneid % 4);
#pragma unroll
for(int j = 0; j < src.width; j++) {
int col = j*src.tile_size_col + (warp_laneid / 4);
dst_ptr[(row+0)*row_stride + (col+0)] = base_types::convertor<U, T>::convert(src.tiles[i][j].data[0].x);
dst_ptr[(row+0)*row_stride + (col+8)] = base_types::convertor<U, T>::convert(src.tiles[i][j].data[1].x);
}
#pragma unroll
for(int j = 0; j < src.width; j++) {
int col = j*src.tile_size_col + (warp_laneid / 4);
dst_ptr[(row+1)*row_stride + (col+0)] = base_types::convertor<U, T>::convert(src.tiles[i][j].data[0].y);
dst_ptr[(row+1)*row_stride + (col+8)] = base_types::convertor<U, T>::convert(src.tiles[i][j].data[1].y);
}
#pragma unroll
for(int j = 0; j < src.width; j++) {
int col = j*src.tile_size_col + (warp_laneid / 4);
dst_ptr[(row+8)*row_stride + (col+0)] = base_types::convertor<U, T>::convert(src.tiles[i][j].data[2].x);
dst_ptr[(row+8)*row_stride + (col+8)] = base_types::convertor<U, T>::convert(src.tiles[i][j].data[3].x);
}
#pragma unroll
for(int j = 0; j < src.width; j++) {
int col = j*src.tile_size_col + (warp_laneid / 4);
dst_ptr[(row+9)*row_stride + (col+0)] = base_types::convertor<U, T>::convert(src.tiles[i][j].data[2].y);
dst_ptr[(row+9)*row_stride + (col+8)] = base_types::convertor<U, T>::convert(src.tiles[i][j].data[3].y);
}
}
}
template<ducks::rt::all RT, ducks::gl::all GL, ducks::coord::tile COORD=coord<rt<typename RT::T, GROUP_WARPS*RT::rows, RT::cols, typename RT::layout>>>
__device__ inline static void store(const GL &dst, const RT &src, const COORD &idx) {
store<2>(dst, src, idx);
}
@@ -1,168 +0,0 @@
/**
* @file
* @brief Group (collaborative warp) ops for loading shared tiles from and storing to global memory.
*/
/**
* @brief Loads data from global memory into a shared memory tile.
*
* @tparam ST The type of the shared tile.
* @param[out] dst The destination shared memory tile.
* @param[in] src The source global memory array.
* @param[in] idx The coordinate of the tile in the global memory array.
*/
template<int axis, bool assume_aligned, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load(ST &dst, const GL &src, const COORD &idx) {
using T = typename ST::dtype;
const int row_stride = src.template stride<axis>();
// we can handle this many rows each time we run a memcpy_async
constexpr int elem_per_memcpy = sizeof(float4)/sizeof(typename ST::dtype);
constexpr int memcpy_per_row = dst.cols / elem_per_memcpy;
constexpr int total_calls = (dst.height*dst.width * kittens::TILE_ROW_DIM<T>*kittens::TILE_COL_DIM<T> + GROUP_THREADS*elem_per_memcpy-1) / (GROUP_THREADS*elem_per_memcpy); // round up
constexpr int total_rows = dst.height*dst.width;
coord<> unit_coord = idx.template unit_coord<axis, 3>();
typename GL::dtype *src_ptr = (typename GL::dtype*)&src[unit_coord];
uint32_t dst_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst.data[0]));
int laneid = threadIdx.x % GROUP_THREADS;
#pragma unroll
for(int i = 0; i < total_calls; i++) {
int load_idx = i * GROUP_THREADS + laneid;
int row = load_idx / memcpy_per_row;
int col = (load_idx*elem_per_memcpy) % dst.cols;
if constexpr (assume_aligned) {
float4 tmp;
move<float4>::ldg(tmp, (float4*)&src_ptr[row*row_stride + col]);
move<float4>::sts(dst.idx(dst_ptr, {row, col}), tmp);
}
else {
if (row + unit_coord.template dim<axis>() < src.template shape<axis>()) {
float4 tmp;
move<float4>::ldg(tmp, (float4*)&src_ptr[row*row_stride + col]);
move<float4>::sts(dst.idx(dst_ptr, {row, col}), tmp);
}
else {
float4 zeros = {0.f,0.f,0.f,0.f};
move<float4>::sts(dst.idx(dst_ptr, {row, col}), zeros); // use the default value
}
}
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load(ST &dst, const GL &src, const COORD &idx) {
load<2, false, ST, GL, COORD>(dst, src, idx);
}
/**
* @brief Stores data from a shared memory tile into global memory.
*
* @tparam ST The type of the shared tile.
* @param[out] dst The destination global memory array.
* @param[in] src The source shared memory tile.
* @param row_stride[in] The stride between rows in the destination array.
*/
template<int axis, bool assume_aligned, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store(const GL &dst, const ST &src, const COORD &idx) {
using T = typename ST::dtype;
const int row_stride = dst.template stride<axis>();
// we can handle this many rows each time we run a memcpy_async
constexpr int elem_per_memcpy = sizeof(float4)/sizeof(typename ST::dtype);
constexpr int memcpy_per_row = src.cols / elem_per_memcpy;
constexpr int total_calls = (src.height*src.width * kittens::TILE_ROW_DIM<T>*kittens::TILE_COL_DIM<T> + GROUP_THREADS*elem_per_memcpy-1) / (GROUP_THREADS*elem_per_memcpy); // round up
coord<> unit_coord = idx.template unit_coord<axis, 3>();
typename GL::dtype *dst_ptr = (typename GL::dtype*)&dst[unit_coord];
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src.data[0]));
int laneid = threadIdx.x % GROUP_THREADS;
#pragma unroll
for(int i = 0; i < total_calls; i++) {
int load_idx = i * GROUP_THREADS + laneid;
int row = load_idx / memcpy_per_row;
int col = (load_idx*elem_per_memcpy) % src.cols;
if constexpr (assume_aligned) {
float4 tmp;
move<float4>::lds(tmp, src.idx(src_ptr, {row, col}));
move<float4>::stg((float4*)&dst_ptr[row*row_stride + col], tmp);
}
else {
if (row + unit_coord.template dim<axis>() < dst.template shape<axis>()) {
float4 tmp;
move<float4>::lds(tmp, src.idx(src_ptr, {row, col}));
move<float4>::stg((float4*)&dst_ptr[row*row_stride + col], tmp);
}
}
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store(const GL &dst, const ST &src, const COORD &idx) {
store<2, false, ST, GL, COORD>(dst, src, idx);
}
/**
* @brief Asynchronously loads data from global memory into a shared memory tile.
*
* @tparam ST The type of the shared tile.
* @param[out] dst The destination shared memory tile.
* @param[in] src The source global memory array.
*
* @note This function expects 16-byte alignments. Otherwise, behavior is undefined.
*/
template<int axis, bool assume_aligned, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load_async(ST &dst, const GL &src, const COORD &idx) {
using T = typename ST::dtype;
const int row_stride = src.template stride<axis>();
// we can handle this many rows each time we run a memcpy_async
constexpr int elem_per_memcpy = sizeof(float4)/sizeof(typename ST::dtype);
constexpr int memcpy_per_row = dst.cols / elem_per_memcpy;
constexpr int total_calls = (dst.height*dst.width * kittens::TILE_ROW_DIM<T>*kittens::TILE_COL_DIM<T> + GROUP_THREADS*elem_per_memcpy-1) / (GROUP_THREADS*elem_per_memcpy); // round up
coord<> unit_coord = idx.template unit_coord<axis, 3>();
typename GL::dtype *src_ptr = (typename GL::dtype*)&src[unit_coord];
uint32_t dst_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst.data[0]));
int laneid = threadIdx.x % GROUP_THREADS;
#pragma unroll
for(int i = 0; i < total_calls; i++) {
int load_idx = i * GROUP_THREADS + laneid;
int row = load_idx / memcpy_per_row;
int col = (load_idx*elem_per_memcpy) % dst.cols;
if constexpr (assume_aligned) {
asm volatile(
"cp.async.cg.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(dst.idx(dst_ptr, {row, col})), "l"(&src_ptr[row*row_stride + col])
: "memory"
);
}
else {
if (row + unit_coord.template dim<axis>() < src.template shape<axis>()) {
asm volatile(
"cp.async.cg.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(dst.idx(dst_ptr, {row, col})), "l"(&src_ptr[row*row_stride + col])
: "memory"
);
}
else {
// printf("thread %d skipping async load on row %d, col %d\n", threadIdx.x, row + unit_coord.template dim<axis>(), col);
float4 zeros = {0.f,0.f,0.f,0.f};
move<float4>::sts(dst.idx(dst_ptr, {row, col}), zeros); // use the default value
}
}
}
asm volatile("cp.async.commit_group;\n" ::: "memory");
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load_async(ST &dst, const GL &src, const COORD &idx) {
load_async<2, false, ST, GL, COORD>(dst, src, idx);
}
@@ -1,323 +0,0 @@
/**
* @file
* @brief Functions for a warpgroup to collaboratively transfer data directly between shared memory and registers and back.
*/
/**
* @brief Collaboratively load data from a shared tile into register tiles split across a warpgroup.
*
* @tparam RT The register tile type
* @tparam ST The shared tile type
* @param dst[out] The destination register tile.
* @param src[in] The source shared tile.
*/
template<ducks::rt::all RT, ducks::st::all ST>
__device__ inline static void load(RT &dst, const ST &src) {
constexpr int height = ST::height;
constexpr int warp_height = RT::height;
static_assert(height%GROUP_WARPS == 0, "Group load / store requires tile height to be a multiple of GROUP_WARPS.");
static_assert(height%warp_height == 0, "Group load / store requires tile height to be a multiple of the RT height.");
static_assert(ST::width==RT::width, "Group load / store requires tile widths to match.");
int local_warpid;
if constexpr(GROUP_WARPS % 4 == 0) local_warpid = (warpid()/4+(warpid()%4)*(GROUP_WARPS/4));
else local_warpid = warpid();
using T2 = RT::dtype;
using U = ST::dtype;
using T = base_types::packing<T2>::unpacked_type;
using U2 = base_types::packing<U>::packed_type;
int warp_laneid = ::kittens::laneid();
// convert to shared state space
uint32_t shared_addr = static_cast<uint32_t>(__cvta_generic_to_shared(&src.data[0]));
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
if constexpr (sizeof(typename ST::dtype) == 2) {
// handle the row-major layout for 16-bit types
U2 tmp[4];
int row = (local_warpid*warp_height + i)*dst.tile_size_row + (warp_laneid % 16);
int col = j*dst.tile_size_col + (warp_laneid / 16) * 8;
if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row>) {
move<U2>::ldsm4(tmp[0], tmp[1], tmp[2], tmp[3], src.idx(shared_addr, {row, col}));
}
else {
move<U2>::ldsm4t(tmp[0], tmp[2], tmp[1], tmp[3], src.idx(shared_addr, {row, col}));
}
dst.tiles[i][j].data[0] = base_types::convertor<T2, U2>::convert(tmp[0]);
dst.tiles[i][j].data[1] = base_types::convertor<T2, U2>::convert(tmp[1]);
dst.tiles[i][j].data[2] = base_types::convertor<T2, U2>::convert(tmp[2]);
dst.tiles[i][j].data[3] = base_types::convertor<T2, U2>::convert(tmp[3]);
}
else if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row> && sizeof(typename ST::dtype) == 1) {
// handle the row-major layout for 8-bit types
int warp_group_16 = (warp_laneid / 16); // divide each warp into two groups of 16 threads
int lane_in_16 = warp_laneid % 16; // position in group of 16 threads
int row = (local_warpid*warp_height + i)*dst.tile_size_row + (lane_in_16 % 16); // find base row for warp in warpgroup and then distribute the 16 threads in the warp across the rows
int col = j*dst.tile_size_col + warp_group_16 * 16; // find base column and then *16 for second half of the warp
U2 tmp[4];
if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row>) {
move<U2>::ldsm4(tmp[0], tmp[1], tmp[2], tmp[3], src.idx(shared_addr, {row, col}));
}
else {
move<U2>::ldsm4t(tmp[0], tmp[2], tmp[1], tmp[3], src.idx(shared_addr, {row, col}));
}
dst.tiles[i][j].data[0] = base_types::convertor<T2, U2>::convert(tmp[0]);
dst.tiles[i][j].data[1] = base_types::convertor<T2, U2>::convert(tmp[1]);
dst.tiles[i][j].data[2] = base_types::convertor<T2, U2>::convert(tmp[2]);
dst.tiles[i][j].data[3] = base_types::convertor<T2, U2>::convert(tmp[3]);
}
else if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row> && sizeof(typename ST::dtype) == 4) {
// handle the row-major layout for 32-bit types
int row = (local_warpid*warp_height + i)*dst.tile_size_row + (warp_laneid / 4);
int col = j*dst.tile_size_col + 2*(warp_laneid % 4);
if constexpr (ST::rows != ST::underlying_rows || ST::cols != ST::underlying_cols) { // subtile case
row += src.row_offset;
col += src.col_offset;
}
int blit = sizeof(typename ST::dtype) * ((warp_laneid%4) / 2);
U2 tmp[4];
static constexpr int swizzle_repeat = ST::swizzle_bytes * 8;
static constexpr int subtile_cols = ST::swizzle_bytes / sizeof(U);
const int outer_idx = col/subtile_cols;
const uint32_t addr_1 = shared_addr + sizeof(U)*(outer_idx*ST::underlying_rows*subtile_cols + (row+0)*subtile_cols + col%subtile_cols);
const uint32_t addr_2 = shared_addr + sizeof(U)*(outer_idx*ST::underlying_rows*subtile_cols + (row+8)*subtile_cols + col%subtile_cols);
const int swizzle_1 = blit ^ ((addr_1 % swizzle_repeat) >> 7) << 4;
const int swizzle_2 = blit ^ ((addr_2 % swizzle_repeat) >> 7) << 4;
move<U>::lds(tmp[0].x, (addr_1+ 0)^swizzle_1);
move<U>::lds(tmp[0].y, (addr_1+ 4)^swizzle_1);
move<U>::lds(tmp[2].x, (addr_1+32)^swizzle_1);
move<U>::lds(tmp[2].y, (addr_1+36)^swizzle_1);
move<U>::lds(tmp[1].x, (addr_2+ 0)^swizzle_2);
move<U>::lds(tmp[1].y, (addr_2+ 4)^swizzle_2);
move<U>::lds(tmp[3].x, (addr_2+32)^swizzle_2);
move<U>::lds(tmp[3].y, (addr_2+36)^swizzle_2);
dst.tiles[i][j].data[0] = base_types::convertor<T2, U2>::convert(tmp[0]);
dst.tiles[i][j].data[1] = base_types::convertor<T2, U2>::convert(tmp[1]);
dst.tiles[i][j].data[2] = base_types::convertor<T2, U2>::convert(tmp[2]);
dst.tiles[i][j].data[3] = base_types::convertor<T2, U2>::convert(tmp[3]);
if(blit) {
#pragma unroll
for(int k = 0; k < 4; k++) {
dst.tiles[i][j].data[k] = T2{dst.tiles[i][j].data[k].y, dst.tiles[i][j].data[k].x};
}
}
}
else {
// handle the column-major layout
int row = (local_warpid*warp_height + i)*dst.tile_size_row + 2*(warp_laneid % 4);
int col = j*dst.tile_size_col + (warp_laneid / 4);
U2 tmp[4];
move<U>::lds(tmp[0].x, src.idx(shared_addr, {row+0, col+0}));
move<U>::lds(tmp[0].y, src.idx(shared_addr, {row+1, col+0}));
move<U>::lds(tmp[1].x, src.idx(shared_addr, {row+0, col+8}));
move<U>::lds(tmp[1].y, src.idx(shared_addr, {row+1, col+8}));
move<U>::lds(tmp[2].x, src.idx(shared_addr, {row+8, col+0}));
move<U>::lds(tmp[2].y, src.idx(shared_addr, {row+9, col+0}));
move<U>::lds(tmp[3].x, src.idx(shared_addr, {row+8, col+8}));
move<U>::lds(tmp[3].y, src.idx(shared_addr, {row+9, col+8}));
dst.tiles[i][j].data[0] = base_types::convertor<T2, U2>::convert(tmp[0]);
dst.tiles[i][j].data[1] = base_types::convertor<T2, U2>::convert(tmp[1]);
dst.tiles[i][j].data[2] = base_types::convertor<T2, U2>::convert(tmp[2]);
dst.tiles[i][j].data[3] = base_types::convertor<T2, U2>::convert(tmp[3]);
}
}
}
}
/**
* @brief Collaboratively store data into a shared tile from register tiles split across a warpgroup.
*
* @tparam RT The register tile type
* @tparam ST The shared tile type
* @param dst[out] The destination shared tile.
* @param src[in] The source register tile.
*/
template<ducks::st::all ST, ducks::rt::all RT>
__device__ inline static void store(ST &dst, const RT &src) {
constexpr int height = ST::height;
constexpr int warp_height = RT::height;
static_assert(height%GROUP_WARPS == 0, "Group load / store requires tile height to be a multiple of GROUP_WARPS.");
static_assert(height%warp_height == 0, "Group load / store requires tile height to be a multiple of the RT height.");
static_assert(ST::width==RT::width, "Group load / store requires tile widths to match.");
int local_warpid;
if constexpr(GROUP_WARPS % 4 == 0) local_warpid = (warpid()/4+(warpid()%4)*(GROUP_WARPS/4));
else local_warpid = warpid();
using T2 = RT::dtype;
using U = ST::dtype;
using T = base_types::packing<T2>::unpacked_type;
using U2 = base_types::packing<U>::packed_type;
int warp_laneid = ::kittens::laneid();
// convert to shared state space
uint32_t shared_addr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst.data[0]));
#pragma unroll
for(int i = 0; i < warp_height; i++) {
#pragma unroll
for(int j = 0; j < src.width; j++) {
if constexpr (sizeof(typename ST::dtype) == 2) {
// handle the row-major layout
U2 tmp[4];
tmp[0] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[0]);
tmp[1] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[1]);
tmp[2] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[2]);
tmp[3] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[3]);
#ifdef KITTENS_HOPPER
int row = (local_warpid*warp_height + i)*src.tile_size_row + (warp_laneid % 16);
int col = j*src.tile_size_col + (warp_laneid / 16) * 8;
if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row>) {
move<U2>::stsm4(dst.idx(shared_addr, {row, col}), tmp[0], tmp[1], tmp[2], tmp[3]);
}
else {
move<U2>::stsm4t(dst.idx(shared_addr, {row, col}), tmp[0], tmp[2], tmp[1], tmp[3]);
}
#else
if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row>) {
int row = (local_warpid*warp_height + i)*src.tile_size_row + (warp_laneid / 4);
int col = j*src.tile_size_col + 2*(warp_laneid % 4);
move<U2>::sts(dst.idx(shared_addr, {row+0, col+0}), tmp[0]);
move<U2>::sts(dst.idx(shared_addr, {row+8, col+0}), tmp[1]);
move<U2>::sts(dst.idx(shared_addr, {row+0, col+8}), tmp[2]);
move<U2>::sts(dst.idx(shared_addr, {row+8, col+8}), tmp[3]);
}
else {
int row = (local_warpid*warp_height + i)*src.tile_size_row + 2*(warp_laneid % 4);
int col = j*src.tile_size_col + (warp_laneid / 4);
move<U>::sts(dst.idx(shared_addr, {row+0, col+0}), tmp[0].x);
move<U>::sts(dst.idx(shared_addr, {row+1, col+0}), tmp[0].y);
move<U>::sts(dst.idx(shared_addr, {row+0, col+8}), tmp[1].x);
move<U>::sts(dst.idx(shared_addr, {row+1, col+8}), tmp[1].y);
move<U>::sts(dst.idx(shared_addr, {row+8, col+0}), tmp[2].x);
move<U>::sts(dst.idx(shared_addr, {row+9, col+0}), tmp[2].y);
move<U>::sts(dst.idx(shared_addr, {row+8, col+8}), tmp[3].x);
move<U>::sts(dst.idx(shared_addr, {row+9, col+8}), tmp[3].y);
}
#endif
}
else if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row> && sizeof(typename ST::dtype) == 1) {
// handle the row-major layout for 8-bit types
int warp_group_16 = (warp_laneid / 16); // divide each warp into two groups of 16 threads
int lane_in_16 = warp_laneid % 16; // position in group of 16 threads
int row = (local_warpid*warp_height + i)*src.tile_size_row + (lane_in_16 % 16); // find base row for warp in warpgroup and then distribute the 16 threads in the warp across the rows
int col = j*src.tile_size_col + warp_group_16 * 16; // find base column and then *16 for second half of the warp
U2 tmp[4];
tmp[0] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[0]);
tmp[1] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[1]);
tmp[2] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[2]);
tmp[3] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[3]);
if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row>) {
move<U2>::stsm4(dst.idx(shared_addr, {row, col}), tmp[0], tmp[1], tmp[2], tmp[3]);
}
else {
move<U2>::stsm4t(dst.idx(shared_addr, {row, col}), tmp[0], tmp[2], tmp[1], tmp[3]);
}
}
else if constexpr (std::is_same_v<typename RT::layout, ducks::rt_layout::row> && sizeof(typename ST::dtype) == 4) {
// handle the row-major layout for 32-bit types
int row = (local_warpid*warp_height + i)*src.tile_size_row + (warp_laneid / 4);
int col = j*src.tile_size_col + 2*(warp_laneid % 4);
if constexpr (ST::rows != ST::underlying_rows || ST::cols != ST::underlying_cols) { // subtile case
row += dst.row_offset;
col += dst.col_offset;
}
int blit = sizeof(typename ST::dtype) * ((warp_laneid%4) / 2);
T2 reg_tmp[4];
if(blit) {
#pragma unroll
for(int k = 0; k < 4; k++) {
reg_tmp[k] = T2{src.tiles[i][j].data[k].y, src.tiles[i][j].data[k].x};
}
}
else {
#pragma unroll
for(int k = 0; k < 4; k++) {
reg_tmp[k] = src.tiles[i][j].data[k];
}
}
U2 tmp[4];
tmp[0] = base_types::convertor<U2, T2>::convert(reg_tmp[0]);
tmp[1] = base_types::convertor<U2, T2>::convert(reg_tmp[1]);
tmp[2] = base_types::convertor<U2, T2>::convert(reg_tmp[2]);
tmp[3] = base_types::convertor<U2, T2>::convert(reg_tmp[3]);
static constexpr int swizzle_repeat = ST::swizzle_bytes * 8;
static constexpr int subtile_cols = ST::swizzle_bytes / sizeof(U);
const int outer_idx = col/subtile_cols;
const uint32_t addr_1 = shared_addr + sizeof(U)*(outer_idx*ST::underlying_rows*subtile_cols + (row+0)*subtile_cols + col%subtile_cols);
const uint32_t addr_2 = shared_addr + sizeof(U)*(outer_idx*ST::underlying_rows*subtile_cols + (row+8)*subtile_cols + col%subtile_cols);
const int swizzle_1 = blit ^ ((addr_1 % swizzle_repeat) >> 7) << 4;
const int swizzle_2 = blit ^ ((addr_2 % swizzle_repeat) >> 7) << 4;
move<U>::sts((addr_1+ 0)^swizzle_1, tmp[0].x);
move<U>::sts((addr_1+ 4)^swizzle_1, tmp[0].y);
move<U>::sts((addr_1+32)^swizzle_1, tmp[2].x);
move<U>::sts((addr_1+36)^swizzle_1, tmp[2].y);
move<U>::sts((addr_2+ 0)^swizzle_2, tmp[1].x);
move<U>::sts((addr_2+ 4)^swizzle_2, tmp[1].y);
move<U>::sts((addr_2+32)^swizzle_2, tmp[3].x);
move<U>::sts((addr_2+36)^swizzle_2, tmp[3].y);
}
else {
// handle the column-major layout
int row = (local_warpid*warp_height + i)*src.tile_size_row + 2*(warp_laneid % 4);
int col = j*src.tile_size_col + (warp_laneid / 4);
U2 tmp[4];
tmp[0] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[0]);
tmp[1] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[1]);
tmp[2] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[2]);
tmp[3] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[3]);
move<U>::sts(dst.idx(shared_addr, {row+0, col+0}), tmp[0].x);
move<U>::sts(dst.idx(shared_addr, {row+1, col+0}), tmp[0].y);
move<U>::sts(dst.idx(shared_addr, {row+0, col+8}), tmp[1].x);
move<U>::sts(dst.idx(shared_addr, {row+1, col+8}), tmp[1].y);
move<U>::sts(dst.idx(shared_addr, {row+8, col+0}), tmp[2].x);
move<U>::sts(dst.idx(shared_addr, {row+9, col+0}), tmp[2].y);
move<U>::sts(dst.idx(shared_addr, {row+8, col+8}), tmp[3].x);
move<U>::sts(dst.idx(shared_addr, {row+9, col+8}), tmp[3].y);
}
}
}
}
// Load and store of vectors from/to shared tiles.
template<ducks::rv::naive_layout RV, ducks::st::all ST>
__device__ inline static auto load(RV &dst, const ST &src, int2 row_col) {
KITTENS_CHECK_WARP;
static_assert(ST::cols>=RV::length, "Shared tile must be at least as wide as the vector.");
using T = RV::T;
using U = ST::T;
int warp_laneid = ::kittens::laneid();
// convert to shared state space
uint32_t shared_addr = static_cast<uint32_t>(__cvta_generic_to_shared(&src.data[0]));
#pragma unroll
for(int col = warp_laneid; col < dst.length; col+=WARP_THREADS) {
U tmp;
move<U>::lds(tmp, src.idx(shared_addr, {row_col.x, row_col.y + col}));
dst.data[col/WARP_THREADS][0] = base_types::convertor<T, U>::convert(tmp);
}
}
template<ducks::rv::naive_layout RV, ducks::st::all ST>
__device__ inline static auto store(ST &dst, const RV &src, int2 row_col) {
KITTENS_CHECK_WARP;
static_assert(ST::cols>=RV::length, "Shared tile must be at least as wide as the vector.");
using T = RV::T;
using U = ST::T;
int warp_laneid = ::kittens::laneid();
// convert to shared state space
uint32_t shared_addr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst.data[0]));
#pragma unroll
for(int col = warp_laneid; col < src.length; col+=WARP_THREADS) {
U tmp = base_types::convertor<U, T>::convert(src.data[col/WARP_THREADS][0]);
move<U>::sts(dst.idx(shared_addr, {row_col.x, row_col.y + col}), tmp);
}
}
@@ -1,325 +0,0 @@
/**
* @file
* @brief Group (collaborative warp) ops for loading tensor tiles into register tiles.
*/
/**
* @brief Load data from a tensor tile into a register tile.
*
* @tparam RT The register tile type
* @tparam TM The tensor memory tile type
* @param dst[out] The destination register tile.
* @param src[in] The source tensor tile.
*/
template<ducks::rt::row_layout RT, ducks::tt::all TM>
__device__ inline static void load_async(RT &dst, const TM &src) {
if constexpr (GROUP_WARPS == 1) {
static_assert(RT::height == TM::height, "register tile and tensor tile must match height");
static_assert(RT::width == TM::width, "register tile and tensor tile must match width");
using T2 = RT::dtype;
using U = typename TM::dtype;
using U2 = base_types::packing<typename TM::dtype>::packed_type;
if constexpr (sizeof(typename TM::dtype) == 1) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
asm volatile(
"tcgen05.ld.sync.aligned.16x128b.x2.pack::16b.b32 {%0, %1, %2, %3}, [%4];\n"
: "=r"(*(uint32_t*) &dst.tiles[i][j].data[0]),
"=r"(*(uint32_t*) &dst.tiles[i][j].data[1]),
"=r"(*(uint32_t*) &dst.tiles[i][j].data[2]),
"=r"(*(uint32_t*) &dst.tiles[i][j].data[3])
: "r"(src.addr + ((i * dst.tile_size_row) << 16) + (j * dst.tile_size_col)/(4/(uint32_t)sizeof(U)))
);
}
}
} else if constexpr (sizeof(typename TM::dtype) == 2) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
asm volatile(
"tcgen05.ld.sync.aligned.16x128b.x2.pack::16b.b32 {%0, %1, %2, %3}, [%4];\n"
: "=r"(*(uint32_t*) &dst.tiles[i][j].data[0]),
"=r"(*(uint32_t*) &dst.tiles[i][j].data[1]),
"=r"(*(uint32_t*) &dst.tiles[i][j].data[2]),
"=r"(*(uint32_t*) &dst.tiles[i][j].data[3])
: "r"(src.addr + ((i * dst.tile_size_row) << 16) + (j * dst.tile_size_col))
);
}
}
}
else if constexpr (sizeof(typename TM::dtype) == 4) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
if constexpr (dst.width%4 == 0) {
#pragma unroll
for(int j = 0; j < dst.width; j+=4) {
U2 data[16];
asm volatile(
"tcgen05.ld.sync.aligned.16x256b.x8.b32 {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, [%32];\n"
: "=f"(data[0].x), "=f"(data[0].y),
"=f"(data[1].x), "=f"(data[1].y),
"=f"(data[2].x), "=f"(data[2].y),
"=f"(data[3].x), "=f"(data[3].y),
"=f"(data[4].x), "=f"(data[4].y),
"=f"(data[5].x), "=f"(data[5].y),
"=f"(data[6].x), "=f"(data[6].y),
"=f"(data[7].x), "=f"(data[7].y),
"=f"(data[8].x), "=f"(data[8].y),
"=f"(data[9].x), "=f"(data[9].y),
"=f"(data[10].x), "=f"(data[10].y),
"=f"(data[11].x), "=f"(data[11].y),
"=f"(data[12].x), "=f"(data[12].y),
"=f"(data[13].x), "=f"(data[13].y),
"=f"(data[14].x), "=f"(data[14].y),
"=f"(data[15].x), "=f"(data[15].y)
: "r"(src.addr + ((i * dst.tile_size_row) << 16) + (j * dst.tile_size_col)/(4/(uint32_t)sizeof(U)))
);
#pragma unroll
for(int k = 0; k < 4; k++) {
dst.tiles[i][j+0].data[k] = base_types::convertor<T2, U2>::convert(data[k]);
dst.tiles[i][j+1].data[k] = base_types::convertor<T2, U2>::convert(data[k+4]);
dst.tiles[i][j+2].data[k] = base_types::convertor<T2, U2>::convert(data[k+8]);
dst.tiles[i][j+3].data[k] = base_types::convertor<T2, U2>::convert(data[k+12]);
}
}
}
else if constexpr (dst.width%2 == 0) {
#pragma unroll
for(int j = 0; j < dst.width; j+=2) {
U2 data[8];
asm volatile(
"tcgen05.ld.sync.aligned.16x256b.x4.b32 {%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, [%16];\n"
: "=f"(data[0].x), "=f"(data[0].y),
"=f"(data[1].x), "=f"(data[1].y),
"=f"(data[2].x), "=f"(data[2].y),
"=f"(data[3].x), "=f"(data[3].y),
"=f"(data[4].x), "=f"(data[4].y),
"=f"(data[5].x), "=f"(data[5].y),
"=f"(data[6].x), "=f"(data[6].y),
"=f"(data[7].x), "=f"(data[7].y)
: "r"(src.addr + ((i * dst.tile_size_row) << 16) + (j * dst.tile_size_col)/(4/(uint32_t)sizeof(U)))
);
#pragma unroll
for(int k = 0; k < 4; k++) {
dst.tiles[i][j+0].data[k] = base_types::convertor<T2, U2>::convert(data[k]);
dst.tiles[i][j+1].data[k] = base_types::convertor<T2, U2>::convert(data[k+4]);
}
}
}
else {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
U2 data[4];
asm volatile(
"tcgen05.ld.sync.aligned.16x256b.x2.b32 {%0, %1, %2, %3, %4, %5, %6, %7}, [%8];\n"
: "=f"(data[0].x), "=f"(data[0].y),
"=f"(data[1].x), "=f"(data[1].y),
"=f"(data[2].x), "=f"(data[2].y),
"=f"(data[3].x), "=f"(data[3].y)
: "r"(src.addr + ((i * dst.tile_size_row) << 16) + (j * dst.tile_size_col)/(4/(uint32_t)sizeof(U)))
);
#pragma unroll
for(int k = 0; k < 4; k++) {
dst.tiles[i][j].data[k] = base_types::convertor<T2, U2>::convert(data[k]);
}
}
}
}
}
}
else {
static_assert(GROUP_WARPS==4 || GROUP_WARPS==8);
constexpr int warp_rows = TM::rows/GROUP_WARPS;
static_assert(TM::cols==RT::cols);
static_assert(warp_rows==RT::rows);
if constexpr (GROUP_WARPS == 4) {
auto src_subtile = src.template subtile<tt<typename TM::dtype, warp_rows, TM::cols>>(32*warpid(), 0);
::kittens::group<1>::load_async(dst, src_subtile);
}
else {
auto src_subtile = src.template subtile<tt<typename TM::dtype, warp_rows, TM::cols>>(32*(warpid()%4)+16*(warpid()/4), 0);
::kittens::group<1>::load_async(dst, src_subtile);
}
}
}
/**
* @brief Store data into a tensor tile from a register tile.
*
* @tparam RT The register tile type
* @tparam TM The tensor memory tile type
* @param dst[out] The destination tensor tile.
* @param src[in] The source register tile.
*/
template<ducks::rt::all RT, ducks::tt::all TM>
__device__ inline static void store_async(TM &dst, const RT &src) {
if constexpr (GROUP_WARPS == 1) {
static_assert(RT::height == TM::height, "register tile and tensor tile must match height");
static_assert(RT::width == TM::width, "register tile and tensor tile must match width");
using T2 = RT::dtype;
using T = base_types::packing<T2>::unpacked_type;
using U = TM::dtype;
using U2 = base_types::packing<U>::packed_type;
if constexpr (sizeof(typename TM::dtype) == 2) {
#pragma unroll
for(int i = 0; i < src.height; i++) {
if constexpr (src.width%4 == 0) {
#pragma unroll
for(int j = 0; j < src.width; j+=4) {
asm volatile(
"tcgen05.st.sync.aligned.16x128b.x8.b32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16};\n"
:: "r"(dst.addr + ((i * src.tile_size_row) << 16) + (j * src.tile_size_col)/(4/(uint32_t)sizeof(U))),
"r"(*(uint32_t*)&src.tiles[i][j+0].data[0]),
"r"(*(uint32_t*)&src.tiles[i][j+0].data[1]),
"r"(*(uint32_t*)&src.tiles[i][j+0].data[2]),
"r"(*(uint32_t*)&src.tiles[i][j+0].data[3]),
"r"(*(uint32_t*)&src.tiles[i][j+1].data[0]),
"r"(*(uint32_t*)&src.tiles[i][j+1].data[1]),
"r"(*(uint32_t*)&src.tiles[i][j+1].data[2]),
"r"(*(uint32_t*)&src.tiles[i][j+1].data[3]),
"r"(*(uint32_t*)&src.tiles[i][j+2].data[0]),
"r"(*(uint32_t*)&src.tiles[i][j+2].data[1]),
"r"(*(uint32_t*)&src.tiles[i][j+2].data[2]),
"r"(*(uint32_t*)&src.tiles[i][j+2].data[3]),
"r"(*(uint32_t*)&src.tiles[i][j+3].data[0]),
"r"(*(uint32_t*)&src.tiles[i][j+3].data[1]),
"r"(*(uint32_t*)&src.tiles[i][j+3].data[2]),
"r"(*(uint32_t*)&src.tiles[i][j+3].data[3])
);
}
}
else if constexpr (src.width%2 == 0) {
#pragma unroll
for(int j = 0; j < src.width; j+=2) {
asm volatile(
"tcgen05.st.sync.aligned.16x128b.x4.b32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8};\n"
:: "r"(dst.addr + ((i * src.tile_size_row) << 16) + (j * src.tile_size_col)/(4/(uint32_t)sizeof(U))),
"r"(*(uint32_t*)&src.tiles[i][j+0].data[0]),
"r"(*(uint32_t*)&src.tiles[i][j+0].data[1]),
"r"(*(uint32_t*)&src.tiles[i][j+0].data[2]),
"r"(*(uint32_t*)&src.tiles[i][j+0].data[3]),
"r"(*(uint32_t*)&src.tiles[i][j+1].data[0]),
"r"(*(uint32_t*)&src.tiles[i][j+1].data[1]),
"r"(*(uint32_t*)&src.tiles[i][j+1].data[2]),
"r"(*(uint32_t*)&src.tiles[i][j+1].data[3])
);
}
}
else {
#pragma unroll
for(int j = 0; j < src.width; j++) {
asm volatile(
"tcgen05.st.sync.aligned.16x128b.x2.b32 [%0], {%1, %2, %3, %4};\n"
:: "r"(dst.addr + ((i * src.tile_size_row) << 16) + (j * src.tile_size_col)/(4/(uint32_t)sizeof(U))),
"r"(*(uint32_t*)&src.tiles[i][j].data[0]),
"r"(*(uint32_t*)&src.tiles[i][j].data[1]),
"r"(*(uint32_t*)&src.tiles[i][j].data[2]),
"r"(*(uint32_t*)&src.tiles[i][j].data[3])
);
}
}
}
}
else if constexpr (sizeof(typename TM::dtype) == 4) {
#pragma unroll
for(int i = 0; i < src.height; i++) {
if constexpr(src.width%4 == 0) {
#pragma unroll
for(int j = 0; j < src.width; j+=4) {
U2 data[16];
#pragma unroll
for(int k = 0; k < 4; k++) {
data[k] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[k]);
data[k+4] = base_types::convertor<U2, T2>::convert(src.tiles[i][j+1].data[k]);
data[k+8] = base_types::convertor<U2, T2>::convert(src.tiles[i][j+2].data[k]);
data[k+12] = base_types::convertor<U2, T2>::convert(src.tiles[i][j+3].data[k]);
}
asm volatile(
"tcgen05.st.sync.aligned.16x256b.x8.b32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32};\n"
:: "r"(dst.addr + ((i * src.tile_size_row) << 16) + (j * src.tile_size_col)/(4/(uint32_t)sizeof(U))),
"f"(data[0].x), "f"(data[0].y),
"f"(data[1].x), "f"(data[1].y),
"f"(data[2].x), "f"(data[2].y),
"f"(data[3].x), "f"(data[3].y),
"f"(data[4].x), "f"(data[4].y),
"f"(data[5].x), "f"(data[5].y),
"f"(data[6].x), "f"(data[6].y),
"f"(data[7].x), "f"(data[7].y),
"f"(data[8].x), "f"(data[8].y),
"f"(data[9].x), "f"(data[9].y),
"f"(data[10].x), "f"(data[10].y),
"f"(data[11].x), "f"(data[11].y),
"f"(data[12].x), "f"(data[12].y),
"f"(data[13].x), "f"(data[13].y),
"f"(data[14].x), "f"(data[14].y),
"f"(data[15].x), "f"(data[15].y)
);
}
}
else if constexpr(src.width%2 == 0) {
#pragma unroll
for(int j = 0; j < src.width; j+=2) {
U2 data[8];
#pragma unroll
for(int k = 0; k < 4; k++) {
data[k] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[k]);
data[k+4] = base_types::convertor<U2, T2>::convert(src.tiles[i][j+1].data[k]);
}
asm volatile(
"tcgen05.st.sync.aligned.16x256b.x4.b32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16};\n"
:: "r"(dst.addr + ((i * src.tile_size_row) << 16) + (j * src.tile_size_col)/(4/(uint32_t)sizeof(U))),
"f"(data[0].x), "f"(data[0].y),
"f"(data[1].x), "f"(data[1].y),
"f"(data[2].x), "f"(data[2].y),
"f"(data[3].x), "f"(data[3].y),
"f"(data[4].x), "f"(data[4].y),
"f"(data[5].x), "f"(data[5].y),
"f"(data[6].x), "f"(data[6].y),
"f"(data[7].x), "f"(data[7].y)
);
}
}
else {
#pragma unroll
for(int j = 0; j < src.width; j++) {
U2 data[4];
#pragma unroll
for(int k = 0; k < 4; k++) {
data[k] = base_types::convertor<U2, T2>::convert(src.tiles[i][j].data[k]);
}
asm volatile(
"tcgen05.st.sync.aligned.16x256b.x2.b32 [%0], {%1, %2, %3, %4, %5, %6, %7, %8};\n"
:: "r"(dst.addr + ((i * src.tile_size_row) << 16) + (j * src.tile_size_col)/(4/(uint32_t)sizeof(U))),
"f"(data[0].x), "f"(data[0].y),
"f"(data[1].x), "f"(data[1].y),
"f"(data[2].x), "f"(data[2].y),
"f"(data[3].x), "f"(data[3].y)
);
}
}
}
}
}
else {
static_assert(GROUP_WARPS==4 || GROUP_WARPS==8);
constexpr int warp_rows = TM::rows/GROUP_WARPS;
static_assert(TM::cols==RT::cols);
static_assert(warp_rows==RT::rows);
if constexpr (GROUP_WARPS == 4) {
auto dst_subtile = dst.template subtile<tt<typename TM::dtype, warp_rows, TM::cols>>(32*warpid(), 0);
::kittens::group<1>::store_async(dst_subtile, src);
}
else {
auto dst_subtile = dst.template subtile<tt<typename TM::dtype, warp_rows, TM::cols>>(32*(warpid()%4)+16*(warpid()/4), 0);
::kittens::group<1>::store_async(dst_subtile, src);
}
}
}
@@ -1,16 +0,0 @@
/**
* @file
* @brief An aggregate header of group memory operations on tiles.
*/
#include "shared_to_register.cuh"
#include "global_to_register.cuh"
#include "global_to_shared.cuh"
#ifdef KITTENS_BLACKWELL
#include "tensor_to_register.cuh"
#endif
#include "complex/complex_shared_to_register.cuh"
#include "complex/complex_global_to_register.cuh"
#include "complex/complex_global_to_shared.cuh"
@@ -1,134 +0,0 @@
/**
* @file
* @brief Functions for a group scope to call tile TMA functions.
*/
template<int axis, cache_policy policy, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void prefetch(ST &dst, const GL &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::prefetch<axis, policy, ST, GL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void prefetch(ST &dst, const GL &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::prefetch<dim::ROW, cache_policy::NORMAL, ST, GL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_async(const GL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_async<axis, policy, ST, GL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_async(const GL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_async<dim::ROW, cache_policy::NORMAL, ST, GL, COORD>(dst, src, idx);
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::pgl::all PGL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_async(const PGL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_async<axis, policy, ST, PGL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::pgl::all PGL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_async(const PGL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_async<dim::ROW, cache_policy::NORMAL, ST, PGL, COORD>(dst, src, idx);
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_add_async(const GL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_add_async<axis, policy, ST, GL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_add_async(const GL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_add_async<dim::ROW, cache_policy::NORMAL, ST, GL, COORD>(dst, src, idx);
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::pgl::all PGL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_add_async(const PGL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_add_async<axis, policy, ST, PGL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::pgl::all PGL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_add_async(const PGL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_add_async<dim::ROW, cache_policy::NORMAL, ST, PGL, COORD>(dst, src, idx);
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_min_async(const GL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_min_async<axis, policy, ST, GL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_min_async(const GL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_min_async<dim::ROW, cache_policy::NORMAL, ST, GL, COORD>(dst, src, idx);
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::pgl::all PGL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_min_async(const PGL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_min_async<axis, policy, ST, PGL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::pgl::all PGL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_min_async(const PGL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_min_async<dim::ROW, cache_policy::NORMAL, ST, PGL, COORD>(dst, src, idx);
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_max_async(const GL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_max_async<axis, policy, ST, GL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_max_async(const GL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_max_async<dim::ROW, cache_policy::NORMAL, ST, GL, COORD>(dst, src, idx);
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::pgl::all PGL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_max_async(const PGL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_max_async<axis, policy, ST, PGL, COORD>(dst, src, idx); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::pgl::all PGL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void store_max_async(const PGL &dst, const ST &src, const COORD &idx) {
if(laneid() == 0) {
::kittens::tma::store_max_async<dim::ROW, cache_policy::NORMAL, ST, PGL, COORD>(dst, src, idx);
}
}
template<int axis, cache_policy policy, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load_async(ST &dst, const GL &src, const COORD &idx, semaphore& bar) {
if(laneid() == 0) {
::kittens::tma::load_async<axis, policy, ST, GL, COORD>(dst, src, idx, bar); // Don't do the mask
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load_async(ST &dst, const GL &src, const COORD &idx, semaphore& bar) {
if(laneid() == 0) {
::kittens::tma::load_async<dim::ROW, cache_policy::NORMAL, ST, GL, COORD>(dst, src, idx, bar);
}
}
@@ -1,33 +0,0 @@
/**
* @file
* @brief Functions for a group scope to call tile TMA cluster functions.
*/
#ifdef KITTENS_BLACKWELL
template<int axis, cache_policy policy, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load_async(ST &dst, const GL &src, const COORD &idx, semaphore& bar, uint16_t cluster_mask, int dst_mbar_cta=-1) {
if(laneid() == 0) {
::kittens::tma::cluster::load_async<axis, policy, ST, GL, COORD>(dst, src, idx, bar, cluster_mask, dst_mbar_cta);
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load_async(ST &dst, const GL &src, const COORD &idx, semaphore& bar, uint16_t cluster_mask, int dst_mbar_cta=-1) {
if(laneid() == 0) {
::kittens::tma::cluster::load_async<dim::ROW, cache_policy::NORMAL, ST, GL, COORD>(dst, src, idx, bar, cluster_mask, dst_mbar_cta);
}
}
#else
template<int axis, cache_policy policy, ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load_async(ST &dst, const GL &src, const COORD &idx, semaphore& bar, uint16_t cluster_mask) {
if(laneid() == 0) {
::kittens::tma::cluster::load_async<axis, policy, ST, GL, COORD>(dst, src, idx, bar, cluster_mask);
}
}
template<ducks::st::all ST, ducks::gl::all GL, ducks::coord::tile COORD=coord<ST>>
__device__ static inline void load_async(ST &dst, const GL &src, const COORD &idx, semaphore& bar, uint16_t cluster_mask) {
if(laneid() == 0) {
::kittens::tma::cluster::load_async<dim::ROW, cache_policy::NORMAL, ST, GL, COORD>(dst, src, idx, bar, cluster_mask);
}
}
#endif
@@ -1,68 +0,0 @@
/**
* @file
* @brief Various utilities for group TMA memory operations.
*/
/* ---------- Barrier functions for async load ---------- */
/**
* @brief Sets the number of bytes expected at the semaphore.
*
* This function sets the number of bytes expected at the semaphore for the first thread in the warp.
* It converts the semaphore pointer to a generic shared memory pointer and uses an inline assembly
* instruction to set the expected number of bytes.
*
* @param semaphore Reference to the semaphore variable.
* @param bytes The number of bytes expected at the semaphore.
*/
__device__ static inline void expect_bytes(semaphore& bar, uint32_t bytes) {
if(laneid() == 0) {
::kittens::tma::expect_bytes(bar, bytes);
}
}
/**
* @brief Sets the number of bytes expected at the semaphore.
*
* This function sets the number of bytes expected at the mbarrier before the transaction arrives.
*/
template<typename T, typename... args>
__device__ static inline void expect(semaphore& bar, const T& _1, const args&... _2) {
expect_bytes(bar, size_bytes<T, args...>);
}
/* ---------- Synchronization functions for async store ---------- */
/**
* @brief Commits previous asynchronous TMA stores to a group and performs them.
*/
__device__ static inline void store_commit_group() {
asm volatile("cp.async.bulk.commit_group;");
}
/**
* @brief Waits for previous committed TMA store groups to complete.
*
* @tparam N The maximum number of remaining TMA store groups. Defaults to 0.
*/
template <int N=0>
__device__ static inline void store_async_wait() {
asm volatile (
"cp.async.bulk.wait_group %0;"
:
: "n"(N)
: "memory"
);
}
/**
* @brief Waits for previous committed TMA store groups to finish reading from shared memory.
*
* @tparam N The maximum number of remaining TMA store groups. Defaults to 0.
*/
template <int N=0>
__device__ static inline void store_async_read_wait() {
asm volatile (
"cp.async.bulk.wait_group.read %0;"
:
: "n"(N)
: "memory"
);
}
@@ -1,90 +0,0 @@
/**
* @brief Waits for the requested semaphore phase, at cluster scope
*
* @param semaphore Reference to the semaphore variable.
* @param kPhaseBit The phase bit used for the semaphore.
*/
__device__ static inline void wait(semaphore& bar, int kPhaseBit) {
void const* const ptr = &bar;
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
asm volatile (
"{\n"
".reg .pred P1;\n"
"LAB_WAIT:\n"
"mbarrier.try_wait.parity.acquire.cluster.shared::cta.b64 P1, [%0], %1;\n"
"@P1 bra.uni DONE;\n"
"bra.uni LAB_WAIT;\n"
"DONE:\n"
"}\n"
:: "r"(mbar_ptr),
"r"(kPhaseBit)
);
}
/**
* @brief Sets the number of bytes expected at the semaphore, assuming a multicast instruction.
*
* This function sets the number of bytes expected at the semaphore for the first thread in the warp.
* It converts the semaphore pointer to a generic shared memory pointer and uses an inline assembly
* instruction to set the expected number of bytes.
*
* It's worth being aware that this function is particularly necessary for multicast loads, and
* distributed shared memory can actually be done with a normal tma::expect followed by wait. See
* the unit tests of dsmem for an example.
*
* @param semaphore Reference to the semaphore variable.
* @param bytes The number of bytes expected at the semaphore.
*/
__device__ static inline void expect_bytes(semaphore& bar, uint32_t bytes, int dst_cta) {
if(laneid() == 0) {
::kittens::tma::cluster::expect_bytes(bar, bytes, dst_cta);
}
}
/**
* @brief Sets the number of bytes expected at the semaphore.
*
* This function sets the number of bytes expected at the semaphore for the first thread in the warp.
* It converts the semaphore pointer to a generic shared memory pointer and uses an inline assembly
* instruction to set the expected number of bytes.
*
* @tparam T The type of the data to be stored at the semaphore.
* @param semaphore Reference to the semaphore variable.
*/
/**
* @brief Sets the number of bytes expected at the semaphore.
*
* This function sets the number of bytes expected at the mbarrier before the transaction arrives.
*/
template<typename T, typename... args>
__device__ static inline void expect(semaphore& bar, int dst_cta, const T& _1, const args&... _2) {
expect_bytes(bar, size_bytes<T, args...>, dst_cta);
}
/**
* @brief Arrives at a semaphore in cluster scope.
*
* Marks a thread arrival at an mbarrier
*
* @param semaphore Reference to the semaphore variable.
* @param kPhaseBit The phase bit used for the semaphore.
*/
__device__ static inline void arrive(semaphore& bar, int dst_cta, uint32_t count=1) {
if(laneid() == 0) {
::kittens::tma::cluster::arrive(bar, dst_cta, count);
}
}
// Generic transfer
__device__ static inline void store_async(void *dst, void *src, int dst_cta, uint32_t size_bytes, semaphore& bar) {
if(laneid() == 0) {
::kittens::tma::cluster::store_async(dst, src, dst_cta, size_bytes, bar);
}
}
// Templated transfer for convenience
template<typename T>
__device__ static inline void store_async(T &dst_, T &src_, int dst_cta, semaphore& bar) {
store_async((void*)&dst_, (void*)&src_, dst_cta, size_bytes<T>, bar);
}
@@ -1,168 +0,0 @@
/**
* @file
* @brief Various utilities for group memory operations.
*/
template<int N=0> __device__ static inline void load_async_wait(int bar_id) { // for completing (non-TMA) async loads
asm volatile("cp.async.wait_group %0;\n" : : "n"(N) : "memory");
sync(bar_id);
}
template<int N=0> __device__ static inline void load_async_wait() { // for completing (non-TMA) async loads
KITTENS_CHECK_WARP
asm volatile("cp.async.wait_group %0;\n" : : "n"(N) : "memory");
__syncwarp();
}
__device__ static inline void arrive(barrier<GROUP_WARPS> bar) {
asm volatile("bar.arrive %0, %1;\n" :: "r"(bar.barrier_id), "n"(GROUP_WARPS*WARP_THREADS) : "memory");
}
__device__ static inline void arrive_and_wait(barrier<GROUP_WARPS> bar) {
asm volatile("bar.sync %0, %1;\n" :: "r"(bar.barrier_id), "n"(GROUP_WARPS*WARP_THREADS) : "memory");
}
/**
* @brief Initializes a synchronization semaphore with a transaction count and sets the expected number of bytes.
*
* This function sets up a semaphore that is used to synchronize threads within a block during asynchronous operations.
* It initializes the semaphore with a thread count semaphore.
*
* Additionally, if it is given a shared tile type, it will also call `set_bytes` to prepare for the memory transaction.
*
* @param[out] semaphore The semaphore variable to initialize.
* @param[in] tc The thread counter for the semaphore.
*/
__device__ static inline void init_semaphore(semaphore& bar, int thread_count, int transaction_count=0) {
if (laneid() == 0) {
void const* const ptr = &bar;
uint32_t bar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
asm volatile (
"mbarrier.init.shared::cta.b64 [%0], %1;\n"
:: "r"(bar_ptr), "r"(thread_count+transaction_count)
);
}
}
/**
* @brief Invalidate an mbarrier
*
* @param[out] semaphore The semaphore variable to initialize.
* @param[in] tc The thread counter for the semaphore.
*/
__device__ static inline void invalidate_semaphore(semaphore& bar) {
if (laneid() == 0) {
void const* const ptr = &bar;
uint32_t bar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
asm volatile (
"mbarrier.inval.shared::cta.b64 [%0];\n"
:: "r"(bar_ptr)
);
}
}
/**
* @brief Arrives at a semaphore.
*
* Marks a warp arrival at an mbarrier
*
* @param semaphore Reference to the semaphore variable.
* @param kPhaseBit The phase bit used for the semaphore.
*/
__device__ static inline void arrive(semaphore& sem) {
if(laneid() == 0) {
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&sem));
asm volatile (
"mbarrier.arrive.release.cta.shared::cta.b64 _, [%0];\n"
:
: "r"(mbar_ptr)
: "memory"
);
}
}
template<int num_warps> __device__ static inline void arrive(barrier<num_warps> bar) {
asm volatile("bar.arrive %0, %1;\n" :: "r"(bar.barrier_id), "n"(num_warps*WARP_THREADS) : "memory");
}
#ifdef KITTENS_HOPPER
/**
* @brief Arrives at a semaphore.
*
* Marks a warp arrival at an mbarrier
*
* @param semaphore Reference to the semaphore variable.
* @param kPhaseBit The phase bit used for the semaphore.
*/
__device__ static inline void arrive(semaphore& sem, uint32_t count) {
if(laneid() == 0) {
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&sem));
asm volatile (
"mbarrier.arrive.release.cta.shared::cta.b64 _, [%0], %1;\n"
:
: "r"(mbar_ptr), "r"(count)
: "memory"
);
}
}
#endif
/**
* @brief Waits for the requested semaphore phase.
*
* @param semaphore Reference to the semaphore variable.
* @param kPhaseBit The phase bit used for the semaphore.
*/
__device__ static inline void wait(semaphore& sem, int kPhaseBit) {
void const* const ptr = &sem;
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
#ifdef KITTENS_HOPPER
asm volatile (
"{\n"
".reg .pred P1;\n"
"LAB_WAIT:\n"
"mbarrier.try_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
"@P1 bra.uni DONE;\n"
"bra.uni LAB_WAIT;\n"
"DONE:\n"
"}\n"
:: "r"(mbar_ptr),
"r"(kPhaseBit)
);
#else
asm volatile (
"{\n"
".reg .pred P1;\n"
"LAB_WAIT:\n"
"mbarrier.test_wait.parity.shared::cta.b64 P1, [%0], %1;\n"
"@P1 bra.uni DONE;\n"
"nanosleep.u32 5;\n" // wait a few nanoseconds on pre-Hopper architectures to save instruction issue slots
"bra.uni LAB_WAIT;\n"
"DONE:\n"
"}\n"
:: "r"(mbar_ptr),
"r"(kPhaseBit)
);
#endif
}
/**
* @brief Checks if the requested semaphore phase is ready.
*
* @param semaphore Reference to the semaphore variable.
* @param kPhaseBit The phase bit used for the semaphore.
*/
__device__ static inline int test_wait(semaphore& sem, int kPhaseBit) {
void const* const ptr = &sem;
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(ptr));
int result;
asm volatile (
"{\n"
".reg .pred P1;\n"
"mbarrier.test_wait.parity.shared::cta.b64 P1, [%1], %2;\n"
"selp.u32 %0,1,0,P1;"
"}\n"
: "=r"(result)
: "r"(mbar_ptr), "r"(kPhaseBit)
);
return result;
}
@@ -1,138 +0,0 @@
/**
* @file
* @brief Functions for a warpgroup to collaboratively transfer data directly between global memory and registers and back.
*/
/**
* @brief Collaboratively loads data into register vectors from a source array in global memory.
*
* @tparam RV The register vector type.
* @tparam U The data type of the source array.
* @param[out] dst The destination register vector to load data into.
* @param[in] src The source array in global memory to load data from.
*/
template<ducks::rv::all RV, ducks::gl::all GL>
__device__ inline static void load(RV &dst, const GL &src, const coord<rv<typename RV::T, GROUP_WARPS*RV::length, typename RV::layout>> &idx) {
if constexpr (GROUP_WARPS == 1) {
using T2 = RV::dtype;
using U = typename GL::dtype;
using U2 = base_types::packing<U>::packed_type;
using T = base_types::packing<T2>::unpacked_type;
U *src_ptr = (U*)&src[(idx.template unit_coord<-1, 3>())];
int laneid = ::kittens::laneid();
if constexpr (std::is_same_v<typename RV::layout, align_l>) {
#pragma unroll
for(auto w = 0; w < (dst.outer_dim+3)/4; w++) {
int idx = w*64 + (laneid/4)*8 + 2*(laneid%4);
int o_dim = w*4 + (laneid/4) / 2;
int i_dim = (laneid/4) % 2;
// this should be a maximally coalesced load.
if(idx < dst.outer_dim*16)
dst[o_dim][i_dim] = base_types::convertor<T2, U2>::convert(*(U2*)&src_ptr[idx]);
}
// now we need to do a bunch of shuffle_sync's to make sure everyone has everything they need.
#pragma unroll
for(auto w = 0; w < dst.outer_dim; w++) {
int leader = 8*(w%4) + (laneid%4); // repeats every 64 columns
dst[w][0] = packed_shfl_sync(MASK_ALL, dst[w][0], leader);
dst[w][1] = packed_shfl_sync(MASK_ALL, dst[w][1], leader+4);
}
}
else if constexpr (std::is_same_v<typename RV::layout, ortho_l>) {
// really hoping https://stackoverflow.com/questions/15029765/is-coalescing-triggered-for-accessing-memory-in-reverse-order is still true
// otherwise there will be some pain :/
#pragma unroll
for(auto w = 0; w < (dst.outer_dim+1)/2; w++) {
int idx = w*32 + (laneid%4)*8 + (laneid/4);
int o_dim = w*2 + (laneid%4) / 2;
// this should be a maximally coalesced load.
if(idx < dst.outer_dim*16) {
T tmp = base_types::convertor<T, U>::convert(src_ptr[idx]);
if(laneid%2==0) dst[o_dim][0].x = tmp;
else dst[o_dim][0].y = tmp;
}
}
// now we need to do a bunch of shuffle_sync's to make sure everyone has everything they need.
#pragma unroll
for(auto w = 0; w < dst.outer_dim; w++) {
int leader = (laneid/4)*4 + 2*(w%2); // repeats every 64 columns
dst[w][0].x = __shfl_sync(MASK_ALL, dst[w][0].x, leader);
dst[w][0].y = __shfl_sync(MASK_ALL, dst[w][0].y, leader+1);
}
}
else if constexpr (std::is_same_v<typename RV::layout, naive_l>) {
#pragma unroll
for(auto w = 0; w < dst.outer_dim; w++) {
if(w < dst.outer_dim-1 || dst.length%32 == 0 || laneid<16) {
dst[w][0] = base_types::convertor<T, U>::convert(src_ptr[w*32 + laneid]);
}
}
}
}
else {
// Call warp level load
::kittens::group<1>::load(dst, src, coord<RV>(idx.b, idx.d, idx.r, idx.c*GROUP_WARPS+warpid()));
}
}
/**
* @brief Collaboratively stores data from register vectors to a destination array in global memory.
*
* @tparam RV The register vector type.
* @tparam U The data type of the destination array.
* @param[out] dst The destination array in global memory to store data into.
* @param[in] src The source register vector to store data from.
*/
template<ducks::rv::all RV, ducks::gl::all GL>
__device__ inline static void store(GL &dst, const RV &src, const coord<rv<typename RV::T, GROUP_WARPS*RV::length, typename RV::layout>> &idx) {
if constexpr (GROUP_WARPS == 1) {
using T2 = RV::dtype;
using U = typename GL::dtype;
using U2 = base_types::packing<U>::packed_type;
using T = base_types::packing<T2>::unpacked_type;
U *dst_ptr = (U*)&dst[(idx.template unit_coord<-1, 3>())];
int laneid = ::kittens::laneid();
if constexpr (std::is_same_v<typename RV::layout, align_l>) {
#pragma unroll
for(auto w = 0; w < (src.outer_dim+3)/4; w++) {
int idx = w*64 + (laneid/4)*8 + 2*(laneid%4);
int o_dim = w*4 + (laneid/4) / 2;
int i_dim = (laneid/4) % 2;
// this should be a maximally coalesced store. I hope!
if(idx < src.outer_dim*16)
*(U2*)&dst_ptr[idx] = base_types::convertor<U2, T2>::convert(src[o_dim][i_dim]);
}
}
else if constexpr (std::is_same_v<typename RV::layout, ortho_l>) {
// really hoping https://stackoverflow.com/questions/15029765/is-coalescing-triggered-for-accessing-memory-in-reverse-order is still true
// otherwise there will be some pain :/
#pragma unroll
for(auto w = 0; w < (src.outer_dim+1)/2; w++) {
int idx = w*32 + (laneid%4)*8 + (laneid/4);
int o_dim = w*2 + (laneid%4) / 2;
// this should be a maximally coalesced load.
if(idx < src.outer_dim*16) {
U tmp;
if(laneid%2==0) tmp = base_types::convertor<U, T>::convert(src[o_dim][0].x);
else tmp = base_types::convertor<U, T>::convert(src[o_dim][0].y);
dst_ptr[idx] = tmp;
}
}
}
else if constexpr (std::is_same_v<typename RV::layout, naive_l>) {
#pragma unroll
for(auto w = 0; w < src.outer_dim; w++) {
if(w < src.outer_dim-1 || src.length%32 == 0 || laneid<16) {
dst_ptr[w*32 + laneid] = base_types::convertor<U, T>::convert(src[w][0]);
}
}
}
}
else {
// Call warp level store
::kittens::group<1>::store(dst, src, coord<RV>(idx.b, idx.d, idx.r, idx.c*GROUP_WARPS+warpid()));
}
}
@@ -1,77 +0,0 @@
/**
* @file
* @brief Group (collaborative warp) ops for loading shared vectors from and storing to global memory.
*/
/**
* @brief Loads data from global memory into shared memory vector.
*
* This function loads data from a global memory location pointed to by `src` into a shared memory vector `dst`.
* It calculates the number of elements that can be transferred in one operation based on the size ratio of `float4` to the data type of `SV`.
* The function ensures coalesced memory access and efficient use of bandwidth by dividing the work among threads in a warp.
*
* @tparam SV Shared vector type, must satisfy ducks::sv::all concept.
* @param dst Reference to the shared vector where the data will be loaded.
* @param src Pointer to the global memory location from where the data will be loaded.
*/
template<ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void load(SV &dst, const GL &src, const COORD &idx) {
constexpr uint32_t elem_per_transfer = sizeof(float4) / sizeof(typename SV::dtype);
constexpr uint32_t total_calls = SV::length / elem_per_transfer; // guaranteed to divide
typename GL::dtype *src_ptr = (typename GL::dtype*)&src[(idx.template unit_coord<-1, 3>())];
uint32_t dst_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst.data[0]));
#pragma unroll
for(uint32_t i = threadIdx.x%GROUP_THREADS; i < total_calls; i+=GROUP_THREADS) {
if(i * elem_per_transfer < dst.length) {
float4 tmp;
move<float4>::ldg(tmp, (float4*)&src_ptr[i*elem_per_transfer]);
move<float4>::sts(dst_ptr + sizeof(typename SV::dtype)*i*elem_per_transfer, tmp);
}
}
}
/**
* @brief Stores data from a shared memory vector to global memory.
*
* This function stores data from a shared memory vector `src` to a global memory location pointed to by `dst`.
* Similar to the load function, it calculates the number of elements that can be transferred in one operation based on the size ratio of `float4` to the data type of `SV`.
* The function ensures coalesced memory access and efficient use of bandwidth by dividing the work among threads in a warp.
*
* @tparam SV Shared vector type, must satisfy ducks::sv::all concept.
* @param dst Pointer to the global memory location where the data will be stored.
* @param src Reference to the shared vector from where the data will be stored.
*/
template<ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store(GL &dst, const SV &src, const COORD &idx) {
constexpr uint32_t elem_per_transfer = sizeof(float4) / sizeof(typename SV::dtype);
constexpr uint32_t total_calls = SV::length / elem_per_transfer; // guaranteed to divide
typename GL::dtype *dst_ptr = (typename GL::dtype*)&dst[(idx.template unit_coord<-1, 3>())];
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src.data[0]));
#pragma unroll
for(uint32_t i = threadIdx.x%GROUP_THREADS; i < total_calls; i+=GROUP_THREADS) {
if(i * elem_per_transfer < src.length) {
float4 tmp;
move<float4>::lds(tmp, src_ptr + sizeof(typename SV::dtype)*i*elem_per_transfer);
move<float4>::stg((float4*)&dst_ptr[i*elem_per_transfer], tmp);
}
}
}
template<ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void load_async(SV &dst, const GL &src, const COORD &idx) {
constexpr uint32_t elem_per_transfer = sizeof(float4) / sizeof(typename SV::dtype);
constexpr uint32_t total_calls = SV::length / elem_per_transfer; // guaranteed to divide
typename GL::dtype *src_ptr = (typename GL::dtype*)&src[(idx.template unit_coord<-1, 3>())];
uint32_t dst_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst.data[0]));
#pragma unroll
for(uint32_t i = threadIdx.x%GROUP_THREADS; i < total_calls; i+=GROUP_THREADS) {
if(i * elem_per_transfer < dst.length) {
asm volatile(
"cp.async.cg.shared.global.L2::128B [%0], [%1], 16;\n"
:: "r"(dst_ptr + (uint32_t)sizeof(typename SV::dtype)*i*elem_per_transfer), "l"((uint64_t)&src_ptr[i*elem_per_transfer])
: "memory"
);
}
}
asm volatile("cp.async.commit_group;\n" ::: "memory");
}
@@ -1,159 +0,0 @@
/**
* @file
* @brief Functions for a group to collaboratively transfer data directly between shared memory and registers and back.
*/
/**
* @brief Collaboratively load data from a shared vector into register vectors split across a warpgroup.
*
* @tparam RV The register vector type
* @tparam SV The shared vector type
* @param dst[out] The destination register vector.
* @param src[in] The source shared vector.
*/
template<ducks::rv::all RV, ducks::sv::all SV>
__device__ inline static void load(RV &dst, const SV &src) {
using T2 = RV::dtype;
using U = SV::dtype;
using U2 = base_types::packing<U>::packed_type;
using T = base_types::packing<T2>::unpacked_type;
if constexpr (GROUP_WARPS == 1) {
static_assert(SV::length == RV::length);
int laneid = ::kittens::laneid();
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src.data[0]));
__syncwarp();
if constexpr (std::is_same_v<typename RV::layout, align_l>) {
#pragma unroll
for(auto w = 0; w < (dst.outer_dim+3)/4; w++) {
int idx = w*64 + (laneid/4)*8 + 2*(laneid%4);
int o_dim = w*4 + (laneid/4) / 2;
int i_dim = (laneid/4) % 2;
// this should be a maximally coalesced load.
if(idx < dst.outer_dim*16) {
U2 tmp;
move<U2>::lds(tmp, src_ptr + sizeof(typename SV::dtype)*idx);
dst[o_dim][i_dim] = base_types::convertor<T2, U2>::convert(tmp);
}
}
__syncwarp();
// now we need to do a bunch of shuffle_sync's to make sure everyone has everything they need.
#pragma unroll
for(auto w = 0; w < dst.outer_dim; w++) {
int leader = 8*(w%4) + (laneid%4); // repeats every 64 columns
dst[w][0] = packed_shfl_sync(MASK_ALL, dst[w][0], leader);
dst[w][1] = packed_shfl_sync(MASK_ALL, dst[w][1], leader+4);
}
}
else if constexpr (std::is_same_v<typename RV::layout, ortho_l>) {
// really hoping https://stackoverflow.com/questions/15029765/is-coalescing-triggered-for-accessing-memory-in-reverse-order is still true
// otherwise there will be some pain :/
#pragma unroll
for(auto w = 0; w < (dst.outer_dim+1)/2; w++) {
int idx = w*32 + (laneid%4)*8 + (laneid/4);
int o_dim = w*2 + (laneid%4) / 2;
// this should be a maximally coalesced load.
if(idx < dst.outer_dim*16) {
U tmp;
move<U>::lds(tmp, src_ptr + sizeof(typename SV::dtype)*idx);
if(laneid%2==0) dst[o_dim][0].x = base_types::convertor<T, U>::convert(tmp);
else dst[o_dim][0].y = base_types::convertor<T, U>::convert(tmp);
}
}
__syncwarp();
// now we need to do a bunch of shuffle_sync's to make sure everyone has everything they need.
#pragma unroll
for(auto w = 0; w < dst.outer_dim; w++) {
int leader = (laneid/4)*4 + 2*(w%2); // repeats every 64 columns
dst[w][0].x = __shfl_sync(MASK_ALL, dst[w][0].x, leader);
dst[w][0].y = __shfl_sync(MASK_ALL, dst[w][0].y, leader+1);
}
}
else if constexpr (std::is_same_v<typename RV::layout, naive_l>) {
#pragma unroll
for(auto w = 0; w < dst.outer_dim; w++) {
if(w < dst.outer_dim-1 || RV::length%32 == 0 || laneid<16) {
U tmp;
move<U>::lds(tmp, src_ptr + sizeof(typename SV::dtype)*(w*32 + laneid));
dst[w][0] = base_types::convertor<T, U>::convert(tmp);
}
}
}
}
else {
static_assert(SV::length == RV::length*GROUP_WARPS);// confirm size correct
auto &_src = src.template subvec<RV::length>(warpid()); // pretend it's smaller and do warp-level load
::kittens::group<1>::load(dst, _src); // warp-level
}
}
/**
* @brief Collaboratively store data into a shared vector from register vectors split across a warpgroup.
*
* @tparam RV The register vector type
* @tparam SV The shared vector type
* @param dst[out] The destination shared vector.
* @param src[in] The source register vector.
*/
template<ducks::sv::all SV, ducks::rv::all RV>
__device__ inline static void store(SV &dst, const RV &src) {
using T2 = RV::dtype;
using U = SV::dtype;
using U2 = base_types::packing<U>::packed_type;
using T = base_types::packing<T2>::unpacked_type;
if constexpr (GROUP_WARPS == 1) {
static_assert(SV::length == RV::length);
int laneid = ::kittens::laneid();
uint32_t dst_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst.data[0]));
__syncwarp();
if constexpr (std::is_same_v<typename RV::layout, align_l>) {
#pragma unroll
for(auto w = 0; w < (src.outer_dim+3)/4; w++) {
int idx = w*64 + (laneid/4)*8 + 2*(laneid%4);
int o_dim = w*4 + (laneid/4) / 2;
int i_dim = (laneid/4) % 2;
// this should be a maximally coalesced store. I hope!
if(idx < src.outer_dim*16) {
U2 tmp = base_types::convertor<U2, T2>::convert(src[o_dim][i_dim]);
move<U2>::sts(dst_ptr + sizeof(typename SV::dtype)*idx, tmp);
}
}
}
else if constexpr (std::is_same_v<typename RV::layout, ortho_l>) {
// really hoping https://stackoverflow.com/questions/15029765/is-coalescing-triggered-for-accessing-memory-in-reverse-order is still true
// otherwise there will be some pain :/
#pragma unroll
for(auto w = 0; w < (src.outer_dim+1)/2; w++) {
int idx = w*32 + (laneid%4)*8 + (laneid/4);
int o_dim = w*2 + (laneid%4) / 2;
// this should be a maximally coalesced load.
if(idx < src.outer_dim*16) {
U tmp;
if(laneid%2==0) tmp = base_types::convertor<U, T>::convert(src[o_dim][0].x);
else tmp = base_types::convertor<U, T>::convert(src[o_dim][0].y);
move<U>::sts(dst_ptr + sizeof(typename SV::dtype)*idx, tmp);
}
}
}
else if constexpr (std::is_same_v<typename RV::layout, naive_l>) {
#pragma unroll
for(auto w = 0; w < src.outer_dim; w++) {
if(w < src.outer_dim-1 || RV::length%32 == 0 || laneid<16) {
U tmp = base_types::convertor<U, T>::convert(src[w][0]);
move<U>::sts(dst_ptr + sizeof(typename SV::dtype)*(w*32 + laneid), tmp);
}
}
}
}
else {
static_assert(SV::length == RV::length*GROUP_WARPS);// confirm size correct
auto &_dst = dst.template subvec<RV::length>(warpid()); // pretend it's smaller and do warp-level load
::kittens::group<1>::store(_dst, src); // warp-level
}
}
@@ -1,221 +0,0 @@
/**
* @file
* @brief Functions for a group scope to call vec TMA functions.
*/
/* ---------- Prefetch Tensor Map ---------- */
/**
* @brief Prefetches data from global memory into a shared memory vector, along with the tensormap.
*
* @tparam SV A shared vector type with a TMA-compatible layout
* @param[out] dst The destination shared memory vector.
* @param[in] src_tma_map The source tensormap address in global memory
* @param[in] vec_idx The coord of the requested vector.
*/
template<cache_policy policy, ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void prefetch(SV &dst, const GL &src, const COORD &idx) {
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(src.template get_tma<SV, -1>());
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
::kittens::detail::tma::vec_prefetch_tma_internal<policy>(tma_ptr, tma_coord);
}
}
__KITTENS_TMA_DEFINE_DEFAULT_LOAD_CACHE_VEC__(prefetch)
/* ---------- Async load and store data from gmem/smem ---------- */
/**
* @brief Asynchronously stores data into global memory from a shared memory vector.
*
* This function performs an asynchronous copy operation using CUDA's cp.async.bulk.tensor instruction.
*
* @tparam SV A shared vector type with a TMA-compatible layout
* @param[out] dst_tma_map The destination tensormap address in global memory
* @param[in] src The source shared memory vector.
* @param[in] vec_idx The coord of the vector destination.
*/
template<cache_policy policy, ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store_async(const GL &dst, const SV &src, const COORD &idx) {
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(dst.template get_tma<SV, -1>());
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t src_i_ptr = src_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_store_async_tma_internal<policy>(tma_ptr, src_i_ptr, tma_coord);
}
store_commit_group();
}
__KITTENS_TMA_DEFINE_DEFAULT_STORE_CACHE_VEC__(store_async)
template<cache_policy policy, ducks::sv::all SV, ducks::pgl::all PGL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store_async(const PGL &dst, const SV &src, const COORD &idx) {
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(dst.template get_tma<SV, -1>());
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t src_i_ptr = src_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_store_async_tma_internal<policy>(tma_ptr, src_i_ptr, tma_coord);
}
store_commit_group();
}
__KITTENS_TMA_DEFINE_PGL_DEFAULT_STORE_CACHE_VEC__(store_async)
/**
* @brief Asynchronously performs an add reduction and stores the result into global memory.
*
* This function performs an asynchronous add reduction operation using CUDA's cp.reduce.async.bulk.tensor instruction.
*
* @tparam SV A shared vector type with a TMA-compatible layout
* @param[out] dst_tma_map The destination tensormap address in global memory
* @param[in] src The source shared memory vector.
* @param[in] vec_idx The coord of the vector destination.
*/
template<cache_policy policy, ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store_add_async(const GL &dst, const SV &src, const COORD &idx) {
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(dst.template get_tma<SV, -1>());
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t src_i_ptr = src_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_store_add_async_tma_internal<policy>(tma_ptr, src_i_ptr, tma_coord);
}
store_commit_group();
}
__KITTENS_TMA_DEFINE_DEFAULT_STORE_CACHE_VEC__(store_add_async)
template<cache_policy policy, ducks::sv::all SV, ducks::pgl::all PGL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store_add_async(const PGL &dst, const SV &src, const COORD &idx) {
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(dst.template get_tma<SV, -1>());
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t src_i_ptr = src_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_store_add_async_tma_internal<policy>(tma_ptr, src_i_ptr, tma_coord);
}
store_commit_group();
}
__KITTENS_TMA_DEFINE_PGL_DEFAULT_STORE_CACHE_VEC__(store_add_async)
/**
* @brief Asynchronously performs an min reduction and stores the result into global memory.
*
* This function performs an asynchronous min reduction operation using CUDA's cp.reduce.async.bulk.tensor instruction.
*
* @tparam SV A shared vector type with a TMA-compatible layout
* @param[out] dst_tma_map The destination tensormap address in global memory
* @param[in] src The source shared memory vector.
* @param[in] vec_idx The coord of the vector destination.
*/
template<cache_policy policy, ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store_min_async(const GL &dst, const SV &src, const COORD &idx) {
static_assert(!std::is_same_v<typename SV::dtype, float>, "TMA does not support async min/max reductions for fp32 types.");
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(dst.template get_tma<SV, -1>());
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t src_i_ptr = src_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_store_min_async_tma_internal<policy>(tma_ptr, src_i_ptr, tma_coord);
}
store_commit_group();
}
__KITTENS_TMA_DEFINE_DEFAULT_STORE_CACHE_VEC__(store_min_async)
template<cache_policy policy, ducks::sv::all SV, ducks::pgl::all PGL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store_min_async(const PGL &dst, const SV &src, const COORD &idx) {
static_assert(!std::is_same_v<typename SV::dtype, float>, "TMA does not support async min/max reductions for fp32 types.");
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(dst.template get_tma<SV, -1>());
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t src_i_ptr = src_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_store_min_async_tma_internal<policy>(tma_ptr, src_i_ptr, tma_coord);
}
store_commit_group();
}
__KITTENS_TMA_DEFINE_PGL_DEFAULT_STORE_CACHE_VEC__(store_min_async)
/**
* @brief Asynchronously performs an max reduction and stores the result into global memory.
*
* This function performs an asynchronous max reduction operation using CUDA's cp.reduce.async.bulk.tensor instruction.
*
* @tparam SV A shared vector type with a TMA-compatible layout
* @param[out] dst_tma_map The destination tensormap address in global memory
* @param[in] src The source shared memory vector.
* @param[in] vec_idx The coord of the vector destination.
*/
template<cache_policy policy, ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store_max_async(const GL &dst, const SV &src, const COORD &idx) {
static_assert(!std::is_same_v<typename SV::dtype, float>, "TMA does not support async min/max reductions for fp32 types.");
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(dst.template get_tma<SV, -1>());
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t src_i_ptr = src_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_store_max_async_tma_internal<policy>(tma_ptr, src_i_ptr, tma_coord);
}
store_commit_group();
}
__KITTENS_TMA_DEFINE_DEFAULT_STORE_CACHE_VEC__(store_max_async)
template<cache_policy policy, ducks::sv::all SV, ducks::pgl::all PGL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void store_max_async(const PGL &dst, const SV &src, const COORD &idx) {
static_assert(!std::is_same_v<typename SV::dtype, float>, "TMA does not support async min/max reductions for fp32 types.");
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(dst.template get_tma<SV, -1>());
uint32_t src_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&src));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t src_i_ptr = src_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_store_max_async_tma_internal<policy>(tma_ptr, src_i_ptr, tma_coord);
}
store_commit_group();
}
__KITTENS_TMA_DEFINE_PGL_DEFAULT_STORE_CACHE_VEC__(store_max_async)
/**
* @brief Asynchronously loads data from global memory into a shared memory vector.
*
* This function performs an asynchronous copy operation using CUDA's cp.async.bulk.tensor instruction.
*
* @tparam SV A shared vector type with a TMA-compatible layout
* @param[out] dst The destination shared memory vector.
* @param[in] src_tma_map The source tensormap address in global memory
* @param[in] vec_idx The coord of the requested vector.
* @param[in,out] bar The semaphore used for synchronization of the asynchronous copy.
*/
template<cache_policy policy, ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void load_async(SV &dst, const GL &src, const COORD &idx, semaphore& bar) {
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(src.template get_tma<SV, -1>());
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&bar));
uint32_t dst_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t dst_i_ptr = dst_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::vec_load_async_tma_internal<policy>(tma_ptr, dst_i_ptr, mbar_ptr, tma_coord);
}
}
__KITTENS_TMA_DEFINE_SEMAPHORE_CACHE_VEC__(load_async)
@@ -1,31 +0,0 @@
/**
* @file
* @brief Functions for a group scope to call vec TMA cluster functions.
*/
/**
* @brief Asynchronously loads data from global memory into a shared memory vector, broadcast across a cluster
*
* This function performs an asynchronous copy operation using CUDA's cp.async.bulk.tensor instruction.
*
* @tparam SV A shared vector type with a TMA-compatible layout
* @param[out] dst The destination shared memory vector.
* @param[in] src_tma_map The source tensormap address in global memory
* @param[in,out] bar The semaphore used for synchronization of the asynchronous copy.
* @param[in] vec_idx The coord of the requested vector.
* @param[in] cluster_mask The mask of the clusters to broadcast to.
*/
template<cache_policy policy, ducks::sv::all SV, ducks::gl::all GL, ducks::coord::vec COORD=coord<SV>>
__device__ static inline void load_async(SV &dst, const GL &src, const COORD &idx, semaphore& bar, uint16_t cluster_mask, int dst_mbar_cta=-1) {
coord<> unit_coord = idx.template unit_coord<-1, 3>();
uint64_t tma_ptr = reinterpret_cast<uint64_t>(src.template get_tma<SV, -1>());
uint32_t mbar_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&bar));
uint32_t dst_ptr = static_cast<uint32_t>(__cvta_generic_to_shared(&dst));
for(int i = ::kittens::laneid(); i < ::kittens::detail::tma::sv_tma_dim2<SV>; i += WARP_THREADS) {
coord<> tma_coord = unit_coord;
tma_coord.c += i * ::kittens::detail::tma::sv_tma_dim1<SV>;
uint32_t dst_i_ptr = dst_ptr + i*::kittens::detail::tma::sv_tma_dim1<SV>*sizeof(typename SV::dtype);
::kittens::detail::tma::cluster::vec_load_async_tma_internal<policy>(tma_ptr, dst_i_ptr, mbar_ptr, tma_coord, cluster_mask, dst_mbar_cta);
}
}
__KITTENS_TMA_DEFINE_CLUSTER_SEMAPHORE_CACHE_VEC__(load_async)
@@ -1,8 +0,0 @@
/**
* @file
* @brief An aggregate header of group memory operations on vectors.
*/
#include "shared_to_register.cuh"
#include "global_to_register.cuh"
#include "global_to_shared.cuh"
@@ -1,17 +0,0 @@
/**
* @file
* @brief An aggregate header for all group-scope MMA operations.
*/
// All compilation targets can use the warp-scope MMA operations.
#include "warp/warp.cuh"
// Hopper has its own warpgroup-scope MMA operations.
#ifdef KITTENS_HOPPER
#include "warpgroup/warpgroup.cuh"
#endif
// Blackwell has its own tensor-scope MMA operations.
#ifdef KITTENS_BLACKWELL
#include "tensor/tensor.cuh"
#endif
@@ -1,172 +0,0 @@
/**
* @file Group-level tcgen05 MMA operations.
*/
template<int trans_a, int n_trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B, int acc=1, int ncta=1>
__device__ static inline void mma(D &d, const A &a, const B &b, semaphore &sem) {
if(laneid() == 0) ::kittens::mma<trans_a, n_trans_b, D, A, B, acc, ncta>(d, a, b, sem);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B, int acc=1>
__device__ static inline void mma2(D &d, const A &a, const B &b, semaphore &sem) {
mma<trans_a, trans_b, D, A, B, acc, 2>(d, a, b, sem);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm(D &d, const A &a, const B &b, semaphore &sem) {
mma<trans_a, trans_b, D, A, B, 0>(d, a, b, sem);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2(D &d, const A &a, const B &b, semaphore &sem) {
mma2<trans_a, trans_b, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AB(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::N, transpose::N, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AB(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::N, transpose::N, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_ABt(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::N, transpose::T, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_ABt(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::N, transpose::T, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AtB(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::T, transpose::N, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AtB(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::T, transpose::N, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AtBt(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::T, transpose::T, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AtBt(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::T, transpose::T, D, A, B, 1>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AB(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::N, transpose::N, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AB(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::N, transpose::N, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_ABt(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::N, transpose::T, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_ABt(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::N, transpose::T, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AtB(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::T, transpose::N, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AtB(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::T, transpose::N, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AtBt(D &d, const A &a, const B &b, semaphore &sem) {
mma<transpose::T, transpose::T, D, A, B, 0>(d, a, b, sem);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AtBt(D &d, const A &a, const B &b, semaphore &sem) {
mma2<transpose::T, transpose::T, D, A, B, 0>(d, a, b, sem);
}
// no sem versions
template<int trans_a, int n_trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B, int acc=1, int ncta=1>
__device__ static inline void mma(D &d, const A &a, const B &b) {
if(laneid() == 0) ::kittens::mma<trans_a, n_trans_b, D, A, B, acc, ncta>(d, a, b);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B, int acc=1>
__device__ static inline void mma2(D &d, const A &a, const B &b) {
mma<trans_a, trans_b, D, A, B, acc, 2>(d, a, b);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm(D &d, const A &a, const B &b) {
mma<trans_a, trans_b, D, A, B, 0>(d, a, b);
}
template<int trans_a, int trans_b, ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2(D &d, const A &a, const B &b) {
mma2<trans_a, trans_b, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AB(D &d, const A &a, const B &b) {
mma<transpose::N, transpose::N, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AB(D &d, const A &a, const B &b) {
mma2<transpose::N, transpose::N, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_ABt(D &d, const A &a, const B &b) {
mma<transpose::N, transpose::T, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_ABt(D &d, const A &a, const B &b) {
mma2<transpose::N, transpose::T, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AtB(D &d, const A &a, const B &b) {
mma<transpose::T, transpose::N, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AtB(D &d, const A &a, const B &b) {
mma2<transpose::T, transpose::N, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma_AtBt(D &d, const A &a, const B &b) {
mma<transpose::T, transpose::T, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mma2_AtBt(D &d, const A &a, const B &b) {
mma2<transpose::T, transpose::T, D, A, B, 1>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AB(D &d, const A &a, const B &b) {
mma<transpose::N, transpose::N, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AB(D &d, const A &a, const B &b) {
mma2<transpose::N, transpose::N, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_ABt(D &d, const A &a, const B &b) {
mma<transpose::N, transpose::T, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_ABt(D &d, const A &a, const B &b) {
mma2<transpose::N, transpose::T, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AtB(D &d, const A &a, const B &b) {
mma<transpose::T, transpose::N, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AtB(D &d, const A &a, const B &b) {
mma2<transpose::T, transpose::N, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm_AtBt(D &d, const A &a, const B &b) {
mma<transpose::T, transpose::T, D, A, B, 0>(d, a, b);
}
template<ducks::tt::all D, typename A, ducks::st_descriptor::input B>
__device__ static inline void mm2_AtBt(D &d, const A &a, const B &b) {
mma2<transpose::T, transpose::T, D, A, B, 0>(d, a, b);
}
@@ -1,947 +0,0 @@
/**
* @file
* @brief Matrix multiply-accumulate operations for tiles stored in registers.
*/
/**
* @brief Perform the HMMA.16816 operation.
*
* This function performs the half-precision matrix multiply-accumulate operation
* using the `mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32` instruction.
*
* @param[out] d0 The first half of the output float2 accumulator.
* @param[out] d1 The second half of the output float2 accumulator.
* @param[in] a0 The first half of the first input bf16_2 matrix.
* @param[in] a1 The second half of the first input bf16_2 matrix.
* @param[in] a2 The first half of the second input bf16_2 matrix.
* @param[in] a3 The second half of the second input bf16_2 matrix.
* @param[in] b0 The first half of the bf16_2 matrix B.
* @param[in] b1 The second half of the bf16_2 matrix B.
* @param[in] c0 The first half of the float2 accumulator matrix C.
* @param[in] c1 The second half of the float2 accumulator matrix C.
*/
__device__ static inline void hmma16816( float2 &d0, float2 &d1,
const bf16_2 &a0, const bf16_2 &a1, const bf16_2 &a2, const bf16_2 &a3,
const bf16_2 &b0, const bf16_2 &b1,
const float2 &c0, const float2 &c1 ) {
asm volatile(
// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#multiply-and-accumulate-instruction-mma
"mma.sync.aligned.m16n8k16.row.col.f32.bf16.bf16.f32 " \
"{%0, %1, %2, %3}, " \
"{%4, %5, %6, %7}, " \
"{%8, %9}, " \
"{%10, %11, %12, %13};"
// D matrix
: "+f"(d0.x), "+f"(d0.y),
"+f"(d1.x), "+f"(d1.y)
// A matrix
: "r"(*(uint32_t*)(&a0)), "r"(*(uint32_t*)(&a1)),
"r"(*(uint32_t*)(&a2)), "r"(*(uint32_t*)(&a3)),
// B matrix
"r"(*(uint32_t*)(&b0)), "r"(*(uint32_t*)(&b1)),
// C matrix
"f"(c0.x), "f"(c0.y),
"f"(c1.x), "f"(c1.y)
);
}
/**
* @brief Perform the HMMA.16816 operation with inputs as fp16 and fp32 accumulators
*
* This function performs the half-precision matrix multiply-accumulate operation
* using the `mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32` instruction.
*
* @param[out] d0 The first half of the output float2 accumulator.
* @param[out] d1 The second half of the output float2 accumulator.
* @param[in] a0 The first half of the first input half_2 matrix.
* @param[in] a1 The second half of the first input half_2 matrix.
* @param[in] a2 The first half of the second input half_2 matrix.
* @param[in] a3 The second half of the second input half_2 matrix.
* @param[in] b0 The first half of the half_2 matrix B.
* @param[in] b1 The second half of the half_2 matrix B.
* @param[in] c0 The first half of the float2 accumulator matrix C.
* @param[in] c1 The second half of the float2 accumulator matrix C.
*/
__device__ static inline void hmma16816( float2 &d0, float2 &d1,
const half_2 &a0, const half_2 &a1, const half_2 &a2, const half_2 &a3,
const half_2 &b0, const half_2 &b1,
const float2 &c0, const float2 &c1 ) {
asm volatile(
// https://docs.nvidia.com/cuda/parallel-thread-execution/#multiply-and-accumulate-instruction-mma
"mma.sync.aligned.m16n8k16.row.col.f32.f16.f16.f32 " \
"{%0, %1, %2, %3}, " \
"{%4, %5, %6, %7}, " \
"{%8, %9}, " \
"{%10, %11, %12, %13};"
// D matrix
: "+f"(d0.x), "+f"(d0.y),
"+f"(d1.x), "+f"(d1.y)
// A matrix
: "r"(*(uint32_t*)(&a0)), "r"(*(uint32_t*)(&a1)),
"r"(*(uint32_t*)(&a2)), "r"(*(uint32_t*)(&a3)),
// B matrix
"r"(*(uint32_t*)(&b0)), "r"(*(uint32_t*)(&b1)),
// C matrix
"f"(c0.x), "f"(c0.y),
"f"(c1.x), "f"(c1.y)
);
}
/**
* @brief Perform the HMMA.16816 operation.
*
* This function performs the half-precision matrix multiply-accumulate operation
* using the `mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16` instruction.
*
* @param[out] d0 The first half of the output half_2 accumulator.
* @param[out] d1 The second half of the output half_2 accumulator.
* @param[in] a0 The first half of the first input half_2 matrix.
* @param[in] a1 The second half of the first input half_2 matrix.
* @param[in] a2 The first half of the second input half_2 matrix.
* @param[in] a3 The second half of the second input half_2 matrix.
* @param[in] b0 The first half of the half_2 matrix B.
* @param[in] b1 The second half of the half_2 matrix B.
* @param[in] c0 The first half of the half_2 accumulator matrix C.
* @param[in] c1 The second half of the half_2 accumulator matrix C.
*/
__device__ static inline void hmma16816( half_2 &d0, half_2 &d1,
const half_2 &a0, const half_2 &a1, const half_2 &a2, const half_2 &a3,
const half_2 &b0, const half_2 &b1,
const half_2 &c0, const half_2 &c1 ) {
asm volatile(
// https://docs.nvidia.com/cuda/parallel-thread-execution/index.html#multiply-and-accumulate-instruction-mma
"mma.sync.aligned.m16n8k16.row.col.f16.f16.f16.f16 " \
"{%0, %1}, " \
"{%2, %3, %4, %5}, " \
"{%6, %7}, " \
"{%8, %9};"
// D matrix
: "=r"(*(uint32_t*)(&d0)), "=r"(*(uint32_t*)(&d1))
// A matrix
: "r"(*(uint32_t*)(&a0)), "r"(*(uint32_t*)(&a1)),
"r"(*(uint32_t*)(&a2)), "r"(*(uint32_t*)(&a3)),
// B matrix
"r"(*(uint32_t*)(&b0)), "r"(*(uint32_t*)(&b1)),
// C matrix
"r"(*(uint32_t*)(&c0)), "r"(*(uint32_t*)(&c1))
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Perform the HMMA.16816 operation for FP8 using fp8e4m3_2.
*
* Using mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 instruction
* but with fp8e4m3_2 (2 FP8 values) instead of fp8e4m3_4
*/
/**
* @brief Perform the HMMA.16816 operation for FP8.
*
* This function performs the fp8-precision matrix multiply-accumulate operation
* using the `mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32` instruction.
*
* @param[out] d0 The first half of the output float2 accumulator.
* @param[out] d1 The second half of the output float2 accumulator.
* @param[in] a0,a1,a2,a3 Input FP8 matrix A values
* @param[in] b0,b1 Input FP8 matrix B values
* @param[in] c0,c1 Input float2 accumulator matrix C values
*/
__device__ static inline void hmma16816( float2 &d0, float2 &d1,
const fp8e4m3_4 &a0, const fp8e4m3_4 &a1,
const fp8e4m3_4 &a2, const fp8e4m3_4 &a3,
const fp8e4m3_4 &b0, const fp8e4m3_4 &b1,
const float2 &c0, const float2 &c1) {
asm volatile(
"mma.sync.aligned.m16n8k32.row.col.f32.e4m3.e4m3.f32 "
"{%0, %1, %2, %3}, "
"{%4, %5, %6, %7}, "
"{%8, %9}, "
"{%10, %11, %12, %13};"
// D matrix (output)
: "+f"(d0.x), "+f"(d0.y),
"+f"(d1.x), "+f"(d1.y)
// A matrix
: "r"(*(uint32_t*)(&a0)), "r"(*(uint32_t*)(&a1)),
"r"(*(uint32_t*)(&a2)), "r"(*(uint32_t*)(&a3)),
// B matrix
"r"(*(uint32_t*)(&b0)), "r"(*(uint32_t*)(&b1)),
// C matrix
"f"(c0.x), "f"(c0.y),
"f"(c1.x), "f"(c1.y)
);
}
#endif
/**
* @brief Base matrix multiply-accumulate operation for row layout.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<bf16_2, row_layout> matrix.
* @param[in] b The second input rt_base<bf16_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<bf16, ducks::rt_layout::row> &a,
const rt_base<bf16, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
/**
* @brief Base matrix multiply-accumulate operation for row layout
* with fp16 inputs and fp32 accumulators.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, row_layout> matrix.
* @param[in] b The second input rt_base<half_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::row> &a,
const rt_base<half, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Base matrix multiply-accumulate operation for row layout.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<fp8e4m3, row_layout> matrix.
* @param[in] b The second input rt_base<fp8e4m3, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<fp8e4m3, ducks::rt_layout::row> &a,
const rt_base<fp8e4m3, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#endif
/**
* @brief Base matrix multiply-accumulate operation for row layout.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<half_2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, row_layout> matrix.
* @param[in] b The second input rt_base<half_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<half_2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AB_base(rt_base<half, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::row> &a,
const rt_base<half, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<half, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
/**
* @brief Base dot product operation for row layout.
*
* This function performs the base dot product operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<bf16_2, row_layout> matrix.
* @param[in] b The second input rt_base<bf16_2, row_layout> matrix in row-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_ABt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<bf16, ducks::rt_layout::row> &a,
const rt_base<bf16, ducks::rt_layout::row> &b, // in row-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2], // for some reason this one seems to need to be backwards
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3], // for some reason this one seems to need to be backwards
c.data[2], c.data[3]
);
}
/**
* @brief Base dot product operation for row layout
* with fp16 inputs and fp32 accumulators.
*
* This function performs the base dot product operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, row_layout> matrix.
* @param[in] b The second input rt_base<half_2, row_layout> matrix in row-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_ABt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::row> &a,
const rt_base<half, ducks::rt_layout::row> &b, // in row-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2], // for some reason this one seems to need to be backwards
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3], // for some reason this one seems to need to be backwards
c.data[2], c.data[3]
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Base dot product operation for row layout.
*
* This function performs the base dot product operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<fp8e4m3x4, row_layout> matrix.
* @param[in] b The second input rt_base<fp8e4m3x4, row_layout> matrix in row-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_ABt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<fp8e4m3, ducks::rt_layout::row> &a,
const rt_base<fp8e4m3, ducks::rt_layout::row> &b, // in row-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2], // for some reason this one seems to need to be backwards
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3], // for some reason this one seems to need to be backwards
c.data[2], c.data[3]
);
}
#endif
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<bf16_2, col_layout> matrix.
* @param[in] b The second input rt_base<bf16_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<bf16, ducks::rt_layout::col> &a,
const rt_base<bf16, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A
* with fp16 inputs and fp32 accumulators.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, col_layout> matrix.
* @param[in] b The second input rt_base<half_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::col> &a,
const rt_base<half, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<fp8e4m3x4, col_layout> matrix.
* @param[in] b The second input rt_base<fp8e4m3x4, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtB_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<fp8e4m3, ducks::rt_layout::col> &a,
const rt_base<fp8e4m3, ducks::rt_layout::col> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#endif
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A and B.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<bf16_2, col_layout> matrix.
* @param[in] b The second input rt_base<bf16_2, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtBt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<bf16, ducks::rt_layout::col> &a,
const rt_base<bf16, ducks::rt_layout::row> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A and B
* with fp16 inputs and fp32 accumulators.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<half_2, col_layout> matrix.
* @param[in] b The second input rt_base<half_2, row_layout> matrix in row-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtBt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<half, ducks::rt_layout::col> &a,
const rt_base<half, ducks::rt_layout::row> &b, // in row-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#ifdef KITTENS_HOPPER
/**
* @brief Base matrix multiply-accumulate operation for row layout with transposed A and B.
*
* This function performs the base matrix multiply-accumulate operation
* using the `hmma16816` function for matrices in row layout.
*
* @param[out] d The output rt_base<float2, row_layout> accumulator.
* @param[in] a The first input rt_base<fp8e4m3x4, col_layout> matrix.
* @param[in] b The second input rt_base<fp8e4m3x4, col_layout> matrix in column-major mode.
* @param[in] c The input rt_base<float2, row_layout> accumulator matrix.
*/
__device__ static inline void mma_AtBt_base(rt_base<float, ducks::rt_layout::row> &d,
const rt_base<fp8e4m3, ducks::rt_layout::col> &a,
const rt_base<fp8e4m3, ducks::rt_layout::row> &b, // in col-major mode
const rt_base<float, ducks::rt_layout::row> &c) {
hmma16816(
d.data[0], d.data[1],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[0], b.data[2],
c.data[0], c.data[1]
);
hmma16816(
d.data[2], d.data[3],
a.data[0], a.data[1], a.data[2], a.data[3],
b.data[1], b.data[3],
c.data[2], c.data[3]
);
}
#endif
/**
* @brief Matrix multiply-accumulate operation.
*
* This function performs the matrix multiply-accumulate operation
* using the `hmma16816` function.
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_hf<N, M, row_layout> accumulator.
* @param[in] a The first input rt_hf<N, K, row_layout> matrix.
* @param[in] b The second input rt_hf<K, M, col_layout> matrix in column-major mode.
* @param[in] c The input rt_hf<N, M, row_layout> accumulator matrix.
*/
template<ducks::rt::row_layout D, ducks::rt::row_layout A, ducks::rt::col_layout B, ducks::rt::row_layout C>
__device__ static inline void mma_AB(D &d,
const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
static_assert(D::rows == A::rows && D::cols == B::cols); // Check D matches A, B
static_assert(A::cols == B::rows); // Check reduction dim is same
static_assert(D::rows == C::rows && D::cols == C::cols); // Check D matches C
#ifdef KITTENS_HOPPER
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, fp8e4m3> &&
std::is_same_v<typename B::T, fp8e4m3> && std::is_same_v<typename C::T, float>)
);
#else
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>)
);
#endif
#pragma unroll
for(int n = 0; n < D::height; n++) {
#pragma unroll
for(int m = 0; m < D::width; m++) {
mma_AB_base(
d.tiles[n][m],
a.tiles[n][0],
b.tiles[0][m],
c.tiles[n][m]
);
#pragma unroll
for(int k = 1; k < A::width; k++) {
mma_AB_base(
d.tiles[n][m],
a.tiles[n][k],
b.tiles[k][m],
d.tiles[n][m]
);
}
}
}
}
/**
* @brief Dot product operation for row layout.
*
* This function performs the dot product operation
* using the `hmma16816` function.
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_fl<N, M, row_layout> accumulator.
* @param[in] a The first input rt_bf<N, K, row_layout> matrix.
* @param[in] b The second input rt_bf<M, K, row_layout> matrix in row-major mode.
* @param[in] c The input rt_fl<N, M, row_layout> accumulator matrix.
*/
template<ducks::rt::row_layout D, ducks::rt::row_layout A, ducks::rt::row_layout B, ducks::rt::row_layout C>
__device__ static inline void mma_ABt(D &d,
const A &a,
const B &b, // notice row and (M, K) instead of col and (K, M)
const C &c) {
KITTENS_CHECK_WARP
static_assert(D::rows == A::rows && D::cols == B::rows); // Check D matches A, B
static_assert(A::cols == B::cols); // Check reduction dim is same
static_assert(D::rows == C::rows && D::cols == C::cols); // Check D matches C
#ifdef KITTENS_HOPPER
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, fp8e4m3> &&
std::is_same_v<typename B::T, fp8e4m3> && std::is_same_v<typename C::T, float>)
);
#else
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>)
);
#endif
#pragma unroll
for(int n = 0; n < D::height; n++) {
#pragma unroll
for(int m = 0; m < D::width; m++) {
mma_ABt_base(
d.tiles[n][m],
a.tiles[n][0],
b.tiles[m][0],
c.tiles[n][m]
);
#pragma unroll
for(int k = 1; k < A::width; k++) {
mma_ABt_base(
d.tiles[n][m],
a.tiles[n][k],
b.tiles[m][k],
d.tiles[n][m]
);
}
}
}
}
/**
* @brief Matrix multiply-accumulate operation with transposed A.
*
* This function performs the matrix multiply-accumulate operation
* using the `hmma16816` instruction.
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_fl<N, M, row_layout> accumulator.
* @param[in] a The first input rt_bf<K, N, row_layout> matrix.
* @param[in] b The second input rt_bf<K, M, col_layout> matrix in column-major mode.
* @param[in] c The input rt_fl<N, M, row_layout> accumulator matrix.
*/
template<ducks::rt::row_layout D, ducks::rt::col_layout A, ducks::rt::col_layout B, ducks::rt::row_layout C>
__device__ static inline void mma_AtB(D &d,
const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
static_assert(D::rows == A::cols && D::cols == B::cols); // Check D matches A, B
static_assert(A::rows == B::rows); // Check reduction dim is same
static_assert(D::rows == C::rows && D::cols == C::cols); // Check D matches C
#ifdef KITTENS_HOPPER
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, fp8e4m3> &&
std::is_same_v<typename B::T, fp8e4m3> && std::is_same_v<typename C::T, float>)
);
#else
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>)
);
#endif
#pragma unroll
for(int n = 0; n < D::height; n++) {
#pragma unroll
for(int m = 0; m < D::width; m++) {
mma_AtB_base(
d.tiles[n][m],
a.tiles[0][n],
b.tiles[0][m],
c.tiles[n][m]
);
#pragma unroll
for(int k = 1; k < A::height; k++) {
mma_AtB_base(
d.tiles[n][m],
a.tiles[k][n],
b.tiles[k][m],
d.tiles[n][m]
);
}
}
}
}
/**
* @brief Matrix multiply-accumulate operation with transposed A and B.
*
* This function performs the matrix multiply-accumulate operation
* using the `hmma16816` instruction.
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_fl<N, M, row_layout> accumulator.
* @param[in] a The first input rt_bf<K, N, col_layout> matrix.
* @param[in] b The second input rt_bf<M, K, row_layout> matrix in column-major mode.
* @param[in] c The input rt_fl<N, M, row_layout> accumulator matrix.
*/
template<ducks::rt::row_layout D, ducks::rt::col_layout A, ducks::rt::row_layout B, ducks::rt::row_layout C>
__device__ static inline void mma_AtBt(D &d,
const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
static_assert(D::rows == A::cols && D::cols == B::rows); // Check D matches A, B
static_assert(A::rows == B::cols); // Check reduction dim is same
static_assert(D::rows == C::rows && D::cols == C::cols); // Check D matches C
#ifdef KITTENS_HOPPER
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, fp8e4m3> &&
std::is_same_v<typename B::T, fp8e4m3> && std::is_same_v<typename C::T, float>)
);
#else
static_assert(
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, bf16> &&
std::is_same_v<typename B::T, bf16> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, float> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, float>) ||
(std::is_same_v<typename D::T, half> && std::is_same_v<typename A::T, half> &&
std::is_same_v<typename B::T, half> && std::is_same_v<typename C::T, half>)
);
#endif
#pragma unroll
for(int n = 0; n < D::height; n++) {
#pragma unroll
for(int m = 0; m < D::width; m++) {
mma_AtBt_base(
d.tiles[n][m],
a.tiles[0][n],
b.tiles[m][0],
c.tiles[n][m]
);
#pragma unroll
for(int k = 1; k < A::height; k++) {
mma_AtBt_base(
d.tiles[n][m],
a.tiles[k][n],
b.tiles[m][k],
d.tiles[n][m]
);
}
}
}
}
template<int trans_A, int trans_B, ducks::rt::all D, ducks::rt::all A, ducks::rt::all B, ducks::rt::all C>
__device__ static inline void mma(D &d,
const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
if constexpr(trans_A == transpose::T) {
if constexpr(trans_B == transpose::T) {
mma_AtBt(d, a, b, c);
} else {
mma_AtB(d, a, b, c);
}
} else {
if constexpr(trans_B == transpose::T) {
mma_ABt(d, a, b, c);
} else {
mma_AB(d, a, b, c);
}
}
}
template<int trans_A, int trans_B, ducks::rt::all A, ducks::rt::all B, ducks::rt::all C>
__device__ static inline C mma(const A &a,
const B &b,
const C &c) {
KITTENS_CHECK_WARP
C d;
if constexpr(trans_A == transpose::T) {
if constexpr(trans_B == transpose::T) {
mma_AtBt(d, a, b, c);
} else {
mma_AtB(d, a, b, c);
}
} else {
if constexpr(trans_B == transpose::T) {
mma_ABt(d, a, b, c);
} else {
mma_AB(d, a, b, c);
}
}
return d;
}
// --------------------------------------------------------------------------------------------------------------------
// --------------------------------------------------------------------------------------------------------------------
// -------------------------------------------------- COMPLEX INPUTS --------------------------------------------------
// --------------------------------------------------------------------------------------------------------------------
// --------------------------------------------------------------------------------------------------------------------
/**
* @brief Matrix multiply-accumulate operation for complex tiles
*
* This function calls mma_AB with hf arguments
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_cmplx_hf<N, M, row_layout> accumulator.
* @param[in] a The first input rt_cmplx_hf<N, K, row_layout> matrix.
* @param[in] b The second input rt_cmplx_hf<K, M, col_layout> matrix in column-major mode.
* @param[in] c The input rt_cmplx_hf<N, M, row_layout> accumulator matrix.
*/
template<int N, int K, int M>
__device__ static inline void mma_AB(crt_hf<N, M, ducks::rt_layout::row> &d,
const crt_hf<N, K, ducks::rt_layout::row> &a,
const crt_hf<K, M, ducks::rt_layout::col> &b,
const crt_hf<N, M, ducks::rt_layout::row> &c) {
KITTENS_CHECK_WARP
// Copy data from input accumulate register into output
::kittens::group<1>::copy(d.real, c.real);
::kittens::group<1>::copy(d.imag, c.imag);
// Negative on B matrix so we can use single accum register
rt_hf<N, K, ducks::rt_layout::row> tmp;
// Hex value for -1 in float16
constexpr half factor = std::bit_cast<__half>(uint16_t(0xFB80));
::kittens::group<1>::mul(tmp, a.imag, factor);
mma_AB(d.real, a.real, b.real, d.real);
mma_AB(d.real, tmp, b.imag, d.real);
mma_AB(d.imag, a.real, b.imag, d.imag);
mma_AB(d.imag, a.imag, b.real, d.imag);
}
/**
* @brief Matrix multiply-accumulate operation for complex tiles
*
* This function calls mma_AB with bf16 arguments
*
* @tparam N The number of row tiles.
* @tparam K The number of column tiles for the A matrix and row tiles for the B matrix.
* @tparam M The number of column tiles for the B matrix.
* @param[out] d The output rt_cmplx_fl<N, M, row_layout> accumulator.
* @param[in] a The first input rt_cmplx_bf<N, K, row_layout> matrix.
* @param[in] b The second input rt_cmplx_bf<K, M, col_layout> matrix in column-major mode.
* @param[in] c The input rt_cmplx_fl<N, M, row_layout> accumulator matrix.
*/
template<int N, int K, int M>
__device__ static inline void mma_AB(crt_fl<N, M, ducks::rt_layout::row> &d,
const crt_bf<N, K, ducks::rt_layout::row> &a,
const crt_bf<K, M, ducks::rt_layout::col> &b,
const crt_fl<N, M, ducks::rt_layout::row> &c) {
KITTENS_CHECK_WARP
// Copy data from input accumulate register into output
::kittens::group<1>::copy(d.real, c.real);
::kittens::group<1>::copy(d.imag, c.imag);
// Negative on B matrix so we can use single accum register
kittens::rt_bf<N, K, ducks::rt_layout::row> tmp;
// Hex value for -1 in bf16
constexpr bf16 factor = std::bit_cast<__nv_bfloat16>(uint16_t(0xBF80));
::kittens::group<1>::mul(tmp, a.imag, factor);
mma_AB(d.real, a.real, b.real, d.real);
mma_AB(d.real, tmp, b.imag, d.real);
mma_AB(d.imag, a.real, b.imag, d.imag);
mma_AB(d.imag, a.imag, b.real, d.imag);
}
@@ -1,334 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 112, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 112, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %61, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"{%56, %57, %58, %59}, " \
"%60, " \
"p, 1, %63, %62;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %61, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"{%56, %57, %58, %59}, " \
"%60, " \
"p, 1, %63, %62;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %33, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27}, " \
"{%28, %29, %30, %31}, " \
"%32, " \
"p, 1, %35, %34;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 112, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %58, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"%56, " \
"%57, " \
"p, 1, %61, %59, %60;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %58, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"%56, " \
"%57, " \
"p, 1, %61, %59, %60;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %30, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n112k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27}, " \
"%28, " \
"%29, " \
"p, 1, %33, %31, %32;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,813 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 128, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 128, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %69, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"{%64, %65, %66, %67}, " \
"%68, " \
"p, 1, %71, %70;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %69, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"{%64, %65, %66, %67}, " \
"%68, " \
"p, 1, %71, %70;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %39, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %69, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"{%64, %65, %66, %67}, " \
"%68, " \
"p, 1, %70;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %69, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"{%64, %65, %66, %67}, " \
"%68, " \
"p, 1, %70;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 128, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %66, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"%64, " \
"%65, " \
"p, 1, %69, %67, %68;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %66, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"%64, " \
"%65, " \
"p, 1, %69, %67, %68;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %37, %35, %36;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %66, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"%64, " \
"%65, " \
"p, 1, %67;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %66, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63}, " \
"%64, " \
"%65, " \
"p, 1, %67;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y),
"+f"(dst.tiles[0][6].data[0].x), "+f"(dst.tiles[0][6].data[0].y),
"+f"(dst.tiles[0][6].data[1].x), "+f"(dst.tiles[0][6].data[1].y),
"+f"(dst.tiles[0][6].data[2].x), "+f"(dst.tiles[0][6].data[2].y),
"+f"(dst.tiles[0][6].data[3].x), "+f"(dst.tiles[0][6].data[3].y),
"+f"(dst.tiles[0][7].data[0].x), "+f"(dst.tiles[0][7].data[0].y),
"+f"(dst.tiles[0][7].data[1].x), "+f"(dst.tiles[0][7].data[1].y),
"+f"(dst.tiles[0][7].data[2].x), "+f"(dst.tiles[0][7].data[2].y),
"+f"(dst.tiles[0][7].data[3].x), "+f"(dst.tiles[0][7].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %35;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n128k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %35;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][7].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,382 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 144, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 144, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %77, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71}, " \
"{%72, %73, %74, %75}, " \
"%76, " \
"p, 1, %79, %78;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %77, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71}, " \
"{%72, %73, %74, %75}, " \
"%76, " \
"p, 1, %79, %78;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %41, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35}, " \
"{%36, %37, %38, %39}, " \
"%40, " \
"p, 1, %43, %42;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 144, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %74, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71}, " \
"%72, " \
"%73, " \
"p, 1, %77, %75, %76;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %74, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71}, " \
"%72, " \
"%73, " \
"p, 1, %77, %75, %76;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %38, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n144k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35}, " \
"%36, " \
"%37, " \
"p, 1, %41, %39, %40;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,190 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 16, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 16, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %13, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"{%8, %9, %10, %11}, " \
"%12, " \
"p, 1, %15, %14;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %13, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"{%8, %9, %10, %11}, " \
"%12, " \
"p, 1, %15, %14;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %9, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f16.f16.f16 " \
"{%0, %1, %2, %3}, " \
"{%4, %5, %6, %7}, " \
"%8, " \
"p, 1, %11, %10;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 16, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %13, %11, %12;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %13, %11, %12;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %6, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n16k16.f16.f16.f16 " \
"{%0, %1, %2, %3}, " \
"%4, " \
"%5, " \
"p, 1, %9, %7, %8;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,666 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 160, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 160, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %85, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"{%80, %81, %82, %83}, " \
"%84, " \
"p, 1, %87, %86;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %85, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"{%80, %81, %82, %83}, " \
"%84, " \
"p, 1, %87, %86;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %45, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"{%40, %41, %42, %43}, " \
"%44, " \
"p, 1, %47, %46;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %85, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"{%80, %81, %82, %83}, " \
"%84, " \
"p, 1, %86;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %85, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"{%80, %81, %82, %83}, " \
"%84, " \
"p, 1, %86;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 160, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %82, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"%80, " \
"%81, " \
"p, 1, %85, %83, %84;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %82, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"%80, " \
"%81, " \
"p, 1, %85, %83, %84;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %42, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"%40, " \
"%41, " \
"p, 1, %45, %43, %44;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %82, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"%80, " \
"%81, " \
"p, 1, %83;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %82, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n160k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79}, " \
"%80, " \
"%81, " \
"p, 1, %83;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,430 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 176, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 176, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %93, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87}, " \
"{%88, %89, %90, %91}, " \
"%92, " \
"p, 1, %95, %94;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %93, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87}, " \
"{%88, %89, %90, %91}, " \
"%92, " \
"p, 1, %95, %94;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %49, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43}, " \
"{%44, %45, %46, %47}, " \
"%48, " \
"p, 1, %51, %50;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 176, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %90, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87}, " \
"%88, " \
"%89, " \
"p, 1, %93, %91, %92;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %90, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87}, " \
"%88, " \
"%89, " \
"p, 1, %93, %91, %92;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %46, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n176k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43}, " \
"%44, " \
"%45, " \
"p, 1, %49, %47, %48;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,674 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 192, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 192, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %101, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"{%96, %97, %98, %99}, " \
"%100, " \
"p, 1, %103, %102;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %101, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"{%96, %97, %98, %99}, " \
"%100, " \
"p, 1, %103, %102;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %55, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %101, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"{%96, %97, %98, %99}, " \
"%100, " \
"p, 1, %102;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %101, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"{%96, %97, %98, %99}, " \
"%100, " \
"p, 1, %102;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 192, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %98, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"%96, " \
"%97, " \
"p, 1, %101, %99, %100;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %98, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"%96, " \
"%97, " \
"p, 1, %101, %99, %100;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %53, %51, %52;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %98, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n192k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95}, " \
"%96, " \
"%97, " \
"p, 1, %99;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,478 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 208, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 208, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %109, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103}, " \
"{%104, %105, %106, %107}, " \
"%108, " \
"p, 1, %111, %110;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %109, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103}, " \
"{%104, %105, %106, %107}, " \
"%108, " \
"p, 1, %111, %110;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %57, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51}, " \
"{%52, %53, %54, %55}, " \
"%56, " \
"p, 1, %59, %58;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 208, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %106, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103}, " \
"%104, " \
"%105, " \
"p, 1, %109, %107, %108;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %106, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103}, " \
"%104, " \
"%105, " \
"p, 1, %109, %107, %108;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %54, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n208k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51}, " \
"%52, " \
"%53, " \
"p, 1, %57, %55, %56;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,826 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 224, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 224, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %117, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"{%112, %113, %114, %115}, " \
"%116, " \
"p, 1, %119, %118;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %117, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"{%112, %113, %114, %115}, " \
"%116, " \
"p, 1, %119, %118;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %61, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"{%56, %57, %58, %59}, " \
"%60, " \
"p, 1, %63, %62;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %117, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"{%112, %113, %114, %115}, " \
"%116, " \
"p, 1, %118;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %117, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"{%112, %113, %114, %115}, " \
"%116, " \
"p, 1, %118;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 224, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %114, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"%112, " \
"%113, " \
"p, 1, %117, %115, %116;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %114, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"%112, " \
"%113, " \
"p, 1, %117, %115, %116;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %58, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55}, " \
"%56, " \
"%57, " \
"p, 1, %61, %59, %60;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %114, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"%112, " \
"%113, " \
"p, 1, %117, %115, %116;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %114, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n224k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111}, " \
"%112, " \
"%113, " \
"p, 1, %117, %115, %116;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,526 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 240, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 240, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %125, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111, %112, %113, %114, %115, %116, %117, %118, %119}, " \
"{%120, %121, %122, %123}, " \
"%124, " \
"p, 1, %127, %126;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y),
"+f"(dst.tiles[0][14].data[0].x), "+f"(dst.tiles[0][14].data[0].y),
"+f"(dst.tiles[0][14].data[1].x), "+f"(dst.tiles[0][14].data[1].y),
"+f"(dst.tiles[0][14].data[2].x), "+f"(dst.tiles[0][14].data[2].y),
"+f"(dst.tiles[0][14].data[3].x), "+f"(dst.tiles[0][14].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %125, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111, %112, %113, %114, %115, %116, %117, %118, %119}, " \
"{%120, %121, %122, %123}, " \
"%124, " \
"p, 1, %127, %126;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y),
"+f"(dst.tiles[0][14].data[0].x), "+f"(dst.tiles[0][14].data[0].y),
"+f"(dst.tiles[0][14].data[1].x), "+f"(dst.tiles[0][14].data[1].y),
"+f"(dst.tiles[0][14].data[2].x), "+f"(dst.tiles[0][14].data[2].y),
"+f"(dst.tiles[0][14].data[3].x), "+f"(dst.tiles[0][14].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %65, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59}, " \
"{%60, %61, %62, %63}, " \
"%64, " \
"p, 1, %67, %66;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 240, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %122, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111, %112, %113, %114, %115, %116, %117, %118, %119}, " \
"%120, " \
"%121, " \
"p, 1, %125, %123, %124;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y),
"+f"(dst.tiles[0][14].data[0].x), "+f"(dst.tiles[0][14].data[0].y),
"+f"(dst.tiles[0][14].data[1].x), "+f"(dst.tiles[0][14].data[1].y),
"+f"(dst.tiles[0][14].data[2].x), "+f"(dst.tiles[0][14].data[2].y),
"+f"(dst.tiles[0][14].data[3].x), "+f"(dst.tiles[0][14].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %122, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59, %60, %61, %62, %63, %64, %65, %66, %67, %68, %69, %70, %71, %72, %73, %74, %75, %76, %77, %78, %79, %80, %81, %82, %83, %84, %85, %86, %87, %88, %89, %90, %91, %92, %93, %94, %95, %96, %97, %98, %99, %100, %101, %102, %103, %104, %105, %106, %107, %108, %109, %110, %111, %112, %113, %114, %115, %116, %117, %118, %119}, " \
"%120, " \
"%121, " \
"p, 1, %125, %123, %124;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][ 0].data[0].x), "+f"(dst.tiles[0][ 0].data[0].y),
"+f"(dst.tiles[0][ 0].data[1].x), "+f"(dst.tiles[0][ 0].data[1].y),
"+f"(dst.tiles[0][ 0].data[2].x), "+f"(dst.tiles[0][ 0].data[2].y),
"+f"(dst.tiles[0][ 0].data[3].x), "+f"(dst.tiles[0][ 0].data[3].y),
"+f"(dst.tiles[0][ 1].data[0].x), "+f"(dst.tiles[0][ 1].data[0].y),
"+f"(dst.tiles[0][ 1].data[1].x), "+f"(dst.tiles[0][ 1].data[1].y),
"+f"(dst.tiles[0][ 1].data[2].x), "+f"(dst.tiles[0][ 1].data[2].y),
"+f"(dst.tiles[0][ 1].data[3].x), "+f"(dst.tiles[0][ 1].data[3].y),
"+f"(dst.tiles[0][ 2].data[0].x), "+f"(dst.tiles[0][ 2].data[0].y),
"+f"(dst.tiles[0][ 2].data[1].x), "+f"(dst.tiles[0][ 2].data[1].y),
"+f"(dst.tiles[0][ 2].data[2].x), "+f"(dst.tiles[0][ 2].data[2].y),
"+f"(dst.tiles[0][ 2].data[3].x), "+f"(dst.tiles[0][ 2].data[3].y),
"+f"(dst.tiles[0][ 3].data[0].x), "+f"(dst.tiles[0][ 3].data[0].y),
"+f"(dst.tiles[0][ 3].data[1].x), "+f"(dst.tiles[0][ 3].data[1].y),
"+f"(dst.tiles[0][ 3].data[2].x), "+f"(dst.tiles[0][ 3].data[2].y),
"+f"(dst.tiles[0][ 3].data[3].x), "+f"(dst.tiles[0][ 3].data[3].y),
"+f"(dst.tiles[0][ 4].data[0].x), "+f"(dst.tiles[0][ 4].data[0].y),
"+f"(dst.tiles[0][ 4].data[1].x), "+f"(dst.tiles[0][ 4].data[1].y),
"+f"(dst.tiles[0][ 4].data[2].x), "+f"(dst.tiles[0][ 4].data[2].y),
"+f"(dst.tiles[0][ 4].data[3].x), "+f"(dst.tiles[0][ 4].data[3].y),
"+f"(dst.tiles[0][ 5].data[0].x), "+f"(dst.tiles[0][ 5].data[0].y),
"+f"(dst.tiles[0][ 5].data[1].x), "+f"(dst.tiles[0][ 5].data[1].y),
"+f"(dst.tiles[0][ 5].data[2].x), "+f"(dst.tiles[0][ 5].data[2].y),
"+f"(dst.tiles[0][ 5].data[3].x), "+f"(dst.tiles[0][ 5].data[3].y),
"+f"(dst.tiles[0][ 6].data[0].x), "+f"(dst.tiles[0][ 6].data[0].y),
"+f"(dst.tiles[0][ 6].data[1].x), "+f"(dst.tiles[0][ 6].data[1].y),
"+f"(dst.tiles[0][ 6].data[2].x), "+f"(dst.tiles[0][ 6].data[2].y),
"+f"(dst.tiles[0][ 6].data[3].x), "+f"(dst.tiles[0][ 6].data[3].y),
"+f"(dst.tiles[0][ 7].data[0].x), "+f"(dst.tiles[0][ 7].data[0].y),
"+f"(dst.tiles[0][ 7].data[1].x), "+f"(dst.tiles[0][ 7].data[1].y),
"+f"(dst.tiles[0][ 7].data[2].x), "+f"(dst.tiles[0][ 7].data[2].y),
"+f"(dst.tiles[0][ 7].data[3].x), "+f"(dst.tiles[0][ 7].data[3].y),
"+f"(dst.tiles[0][ 8].data[0].x), "+f"(dst.tiles[0][ 8].data[0].y),
"+f"(dst.tiles[0][ 8].data[1].x), "+f"(dst.tiles[0][ 8].data[1].y),
"+f"(dst.tiles[0][ 8].data[2].x), "+f"(dst.tiles[0][ 8].data[2].y),
"+f"(dst.tiles[0][ 8].data[3].x), "+f"(dst.tiles[0][ 8].data[3].y),
"+f"(dst.tiles[0][ 9].data[0].x), "+f"(dst.tiles[0][ 9].data[0].y),
"+f"(dst.tiles[0][ 9].data[1].x), "+f"(dst.tiles[0][ 9].data[1].y),
"+f"(dst.tiles[0][ 9].data[2].x), "+f"(dst.tiles[0][ 9].data[2].y),
"+f"(dst.tiles[0][ 9].data[3].x), "+f"(dst.tiles[0][ 9].data[3].y),
"+f"(dst.tiles[0][10].data[0].x), "+f"(dst.tiles[0][10].data[0].y),
"+f"(dst.tiles[0][10].data[1].x), "+f"(dst.tiles[0][10].data[1].y),
"+f"(dst.tiles[0][10].data[2].x), "+f"(dst.tiles[0][10].data[2].y),
"+f"(dst.tiles[0][10].data[3].x), "+f"(dst.tiles[0][10].data[3].y),
"+f"(dst.tiles[0][11].data[0].x), "+f"(dst.tiles[0][11].data[0].y),
"+f"(dst.tiles[0][11].data[1].x), "+f"(dst.tiles[0][11].data[1].y),
"+f"(dst.tiles[0][11].data[2].x), "+f"(dst.tiles[0][11].data[2].y),
"+f"(dst.tiles[0][11].data[3].x), "+f"(dst.tiles[0][11].data[3].y),
"+f"(dst.tiles[0][12].data[0].x), "+f"(dst.tiles[0][12].data[0].y),
"+f"(dst.tiles[0][12].data[1].x), "+f"(dst.tiles[0][12].data[1].y),
"+f"(dst.tiles[0][12].data[2].x), "+f"(dst.tiles[0][12].data[2].y),
"+f"(dst.tiles[0][12].data[3].x), "+f"(dst.tiles[0][12].data[3].y),
"+f"(dst.tiles[0][13].data[0].x), "+f"(dst.tiles[0][13].data[0].y),
"+f"(dst.tiles[0][13].data[1].x), "+f"(dst.tiles[0][13].data[1].y),
"+f"(dst.tiles[0][13].data[2].x), "+f"(dst.tiles[0][13].data[2].y),
"+f"(dst.tiles[0][13].data[3].x), "+f"(dst.tiles[0][13].data[3].y),
"+f"(dst.tiles[0][14].data[0].x), "+f"(dst.tiles[0][14].data[0].y),
"+f"(dst.tiles[0][14].data[1].x), "+f"(dst.tiles[0][14].data[1].y),
"+f"(dst.tiles[0][14].data[2].x), "+f"(dst.tiles[0][14].data[2].y),
"+f"(dst.tiles[0][14].data[3].x), "+f"(dst.tiles[0][14].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %62, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n240k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47, %48, %49, %50, %51, %52, %53, %54, %55, %56, %57, %58, %59}, " \
"%60, " \
"%61, " \
"p, 1, %65, %63, %64;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][ 0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 5].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 6].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 7].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 8].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][ 9].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][10].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][11].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][12].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][13].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][14].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};
File diff suppressed because it is too large Load Diff
@@ -1,446 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 32, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 32, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %23, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %23, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %13, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"{%8, %9, %10, %11}, " \
"%12, " \
"p, 1, %15, %14;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %13, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"{%8, %9, %10, %11}, " \
"%12, " \
"p, 1, %14;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 32, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %21, %19, %20;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %21, %19, %20;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %13, %11, %12;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %19;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %19;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %11;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %10, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n32k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7}, " \
"%8, " \
"%9, " \
"p, 1, %11;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,238 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 48, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 48, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %31, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %31, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %17, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11}, " \
"{%12, %13, %14, %15}, " \
"%16, " \
"p, 1, %19, %18;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 48, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %29, %27, %28;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %29, %27, %28;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %14, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n48k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11}, " \
"%12, " \
"%13, " \
"p, 1, %17, %15, %16;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,587 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 64, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 64, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %39, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %39, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %23, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %37, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"{%32, %33, %34, %35}, " \
"%36, " \
"p, 1, %38;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %21, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"{%16, %17, %18, %19}, " \
"%20, " \
"p, 1, %22;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 64, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %37, %35, %36;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %37, %35, %36;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %21, %19, %20;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %35;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b), // transpose is not supported for FP8
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %34, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31}, " \
"%32, " \
"%33, " \
"p, 1, %35;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b), // transpose is not supported for FP8
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %19;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %18, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n64k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15}, " \
"%16, " \
"%17, " \
"p, 1, %19;\n" \
"}\n"
// a_mat descriptor, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a, imm-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,286 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 80, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 80, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %45, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"{%40, %41, %42, %43}, " \
"%44, " \
"p, 1, %47, %46;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %45, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"{%40, %41, %42, %43}, " \
"%44, " \
"p, 1, %47, %46;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %25, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19}, " \
"{%20, %21, %22, %23}, " \
"%24, " \
"p, 1, %27, %26;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 80, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %42, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"%40, " \
"%41, " \
"p, 1, %45, %43, %44;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %42, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39}, " \
"%40, " \
"%41, " \
"p, 1, %45, %43, %44;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %22, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n80k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19}, " \
"%20, " \
"%21, " \
"p, 1, %25, %23, %24;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,703 +0,0 @@
template<typename T_D, typename T_AB, int trans_a, int trans_b>
struct base<T_D, T_AB, 96, trans_a, trans_b> {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, 96, ducks::rt_layout::row> &dst,
const rt_base<T_AB, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %55, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %55, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %31, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc), "r"(scale_d), "n"(trans_b), "n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %53, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"{%48, %49, %50, %51}, " \
"%52, " \
"p, 1, %54;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %29, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"{%24, %25, %26, %27}, " \
"%28, " \
"p, 1, %30;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "r"(*(uint32_t*)&a_rt.data[0]), "r"(*(uint32_t*)&a_rt.data[1]),
"r"(*(uint32_t*)&a_rt.data[2]), "r"(*(uint32_t*)&a_rt.data[3]),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_b),
"n"(scale_b)
);
}
}
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, 96, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
) {
static_assert(
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) ||
(std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>),
"Invalid type combination for WGMMA."
);
static_assert(scale_b==1 || scale_b==-1, "Invalid scale B (invert) option");
// ----- BF16,BF16 -> FP32 ----- //
if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, bf16>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f32.bf16.bf16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %53, %51, %52;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f32.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %53, %51, %52;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP16,FP16 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, half>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k16.f16.f16.f16 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %29, %27, %28;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
"n"(trans_a),
"n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f32.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %51;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP32 ----- //
else if constexpr (std::is_same_v<T_D, float> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %50, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f32.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23, %24, %25, %26, %27, %28, %29, %30, %31, %32, %33, %34, %35, %36, %37, %38, %39, %40, %41, %42, %43, %44, %45, %46, %47}, " \
"%48, " \
"%49, " \
"p, 1, %51;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+f"(dst.tiles[0][0].data[0].x), "+f"(dst.tiles[0][0].data[0].y),
"+f"(dst.tiles[0][0].data[1].x), "+f"(dst.tiles[0][0].data[1].y),
"+f"(dst.tiles[0][0].data[2].x), "+f"(dst.tiles[0][0].data[2].y),
"+f"(dst.tiles[0][0].data[3].x), "+f"(dst.tiles[0][0].data[3].y),
"+f"(dst.tiles[0][1].data[0].x), "+f"(dst.tiles[0][1].data[0].y),
"+f"(dst.tiles[0][1].data[1].x), "+f"(dst.tiles[0][1].data[1].y),
"+f"(dst.tiles[0][1].data[2].x), "+f"(dst.tiles[0][1].data[2].y),
"+f"(dst.tiles[0][1].data[3].x), "+f"(dst.tiles[0][1].data[3].y),
"+f"(dst.tiles[0][2].data[0].x), "+f"(dst.tiles[0][2].data[0].y),
"+f"(dst.tiles[0][2].data[1].x), "+f"(dst.tiles[0][2].data[1].y),
"+f"(dst.tiles[0][2].data[2].x), "+f"(dst.tiles[0][2].data[2].y),
"+f"(dst.tiles[0][2].data[3].x), "+f"(dst.tiles[0][2].data[3].y),
"+f"(dst.tiles[0][3].data[0].x), "+f"(dst.tiles[0][3].data[0].y),
"+f"(dst.tiles[0][3].data[1].x), "+f"(dst.tiles[0][3].data[1].y),
"+f"(dst.tiles[0][3].data[2].x), "+f"(dst.tiles[0][3].data[2].y),
"+f"(dst.tiles[0][3].data[3].x), "+f"(dst.tiles[0][3].data[3].y),
"+f"(dst.tiles[0][4].data[0].x), "+f"(dst.tiles[0][4].data[0].y),
"+f"(dst.tiles[0][4].data[1].x), "+f"(dst.tiles[0][4].data[1].y),
"+f"(dst.tiles[0][4].data[2].x), "+f"(dst.tiles[0][4].data[2].y),
"+f"(dst.tiles[0][4].data[3].x), "+f"(dst.tiles[0][4].data[3].y),
"+f"(dst.tiles[0][5].data[0].x), "+f"(dst.tiles[0][5].data[0].y),
"+f"(dst.tiles[0][5].data[1].x), "+f"(dst.tiles[0][5].data[1].y),
"+f"(dst.tiles[0][5].data[2].x), "+f"(dst.tiles[0][5].data[2].y),
"+f"(dst.tiles[0][5].data[3].x), "+f"(dst.tiles[0][5].data[3].y)
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e4m3>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f16.e4m3.e4m3 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %27;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
// ----- FP8,FP8 -> FP16 ----- //
else if constexpr (std::is_same_v<T_D, half> && std::is_same_v<T_AB, fp8e5m2>) {
asm volatile (
"{\n"
".reg .pred p;\n" \
"setp.ne.b32 p, %26, 0;\n" \
"wgmma.mma_async.sync.aligned.m64n96k32.f16.e5m2.e5m2 " \
"{%0, %1, %2, %3, %4, %5, %6, %7, %8, %9, %10, %11, %12, %13, %14, %15, %16, %17, %18, %19, %20, %21, %22, %23}, " \
"%24, " \
"%25, " \
"p, 1, %27;\n" \
"}\n"
// a_regs, b_mat descriptor, scale-d, imm-scale-a, imm-scale-b, im-trans-a im-trans-b
: "+r"(*(uint32_t*)&dst.tiles[0][0].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][0].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][1].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][2].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][3].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][4].data[3]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[0]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[1]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[2]),
"+r"(*(uint32_t*)&dst.tiles[0][5].data[3])
: "l"(a_st_desc),
"l"(b_st_desc),
"r"(scale_d),
// "n"(trans_a),
// "n"(trans_b),
"n"(scale_b)
);
}
}
};
@@ -1,47 +0,0 @@
#pragma once
#include "../../../../../common/common.cuh"
#include "../../../../../types/types.cuh"
namespace kittens {
namespace detail {
namespace wgmma {
// templated wrapper for PTX
template<typename T_D, typename T_AB, int cols, int trans_a, int trans_b, int inv=1>
struct base {
template<int scale_b=1> __device__ static inline void rt_st(
rt<T_D, 16, cols, ducks::rt_layout::row> &dst,
const rt<T_AB, 16, cols, ducks::rt_layout::row> & a_rt,
const uint64_t b_st_desc,
int scale_d = 1
);
template<int scale_b=1> __device__ static inline void st_st(
rt<T_D, 16, cols, ducks::rt_layout::row> &dst,
const uint64_t a_st_desc,
const uint64_t b_st_desc,
int scale_d = 1
);
};
// all the ptx's
#include "64x16.impl"
#include "64x32.impl"
#include "64x48.impl"
#include "64x64.impl"
#include "64x80.impl"
#include "64x96.impl"
#include "64x112.impl"
#include "64x128.impl"
#include "64x144.impl"
#include "64x160.impl"
#include "64x176.impl"
#include "64x192.impl"
#include "64x208.impl"
#include "64x224.impl"
#include "64x240.impl"
#include "64x256.impl"
} // namespace wgmma
} // namespace detail
} // namespace kittens
File diff suppressed because it is too large Load Diff
@@ -1,7 +0,0 @@
/**
* @file
* @brief An aggregate header for warp operations on data stored in registers.
*/
#include "tile/tile.cuh"
#include "vec/vec.cuh"
@@ -1,98 +0,0 @@
/**
* @file
* @brief Conversions between data layouts and types for complex register tiles.
*/
/* ---------- LAYOUT SWAPS ---------- */
/**
* @brief Swaps the layout of a complex register tile.
*
* This function swaps the layout of a complex register tile by
* swapping the real and imaginary component tiles' layouts
*
* @tparam T2 The data type of the register tile elements.
* @tparam _height The height of the register tile.
* @tparam _width The width of the register tile.
* @tparam layout The current layout of the register tile.
* @param dst[out] Reference to the destination register tile where the result will be stored.
* @param src[in] Reference to the source register tile to be swapped.
*/
template<typename T2, int _height, int _width, ducks::rt_layout::all layout>
__device__ static inline void swap_layout(crt<T2, _height, _width, typename ducks::rt_layout::transpose<layout>::type> &dst, const crt<T2, _height, _width, layout> &src) {
swap_layout(dst.real, src.real);
swap_layout(dst.real, src.real);
}
/**
* @brief Swaps the layout of a complex register tile in place.
*
* @tparam T2 The data type of the register tile elements.
* @tparam _height The height of the register tile.
* @tparam _width The width of the register tile.
* @tparam layout The current layout of the register tile.
* @param tile[in,out] Reference to the register tile to be swapped in place.
* @return A reference to the swapped register tile.
*/
template<typename T2, int _height, int _width, ducks::rt_layout::all layout>
__device__ static inline crt<T2, _height, _width, typename ducks::rt_layout::transpose<layout>::type>& swap_layout_inplace(crt<T2, _height, _width, layout> &tile) {
tile.real = swap_layout_inplace(tile.real);
tile.imag = swap_layout_inplace(tile.imag);
return tile;
}
/* ---------- TRANSPOSE ---------- */
/**
* @brief Transposes a complex register tile.
*
* This function is marked "sep", which means that the registers underlying dst MUST be separate
* from the registers underlying src.
*
* @tparam T2 The data type of the register tile elements.
* @tparam _height The height of the src register tile, and the width of the dst tile.
* @tparam _width The width of the src register tile, and the height of the dst tile.
* @tparam layout The layout of the register tile.
* @param dst[out] Reference to the register tile in which to store the transposed src.
* @param src[in] Reference to the register tile to be transposed.
*/
template<typename T2, int _height, int _width, ducks::rt_layout::all layout>
__device__ static inline void transpose_sep(crt<T2, _width, _height, layout> &dst, const crt<T2, _height, _width, layout> &src) {
transpose_sep(dst.real, src.real);
transpose_sep(dst.imag, src.imag);
}
/**
* @brief Transposes a square complex register tile in-place.
*
* @tparam T2 The data type of the register tile elements.
* @tparam _height The height (in units of 16) of the src register tile, and the width of the dst tile. (Must be the same as _width.)
* @tparam _width The width (in units of 16) of the src register tile, and the height of the dst tile. (Must be the same as _height.)
* @tparam layout The current layout of the register tile.
* @param src[in] Reference to the register tile to be transposed.
* @return A reference to the transposed register tile.
*/
template<typename T2, int _height, int _width, ducks::rt_layout::all layout>
__device__ static inline crt<T2, _height, _width, layout>& transpose_inplace(crt<T2, _height, _width, layout> &tile) {
tile.real = transpose_inplace(tile.real);
tile.imag = transpose_inplace(tile.imag);
return tile;
}
/* ---------- TYPE SWAPS ---------- */
/**
* @brief Copies a complex register tile, converting the underlying type if necessary.
*
* @tparam T2 The data type of the destination register elements.
* @tparam U2 The data type of the source register elements.
* @tparam _height The height (in units of 16) of the register tiles.
* @tparam _width The width (in units of 16) of the register tiles.
* @tparam layout The current layout of the register tile.
* @param[out] dst A reference to the destination register tile.
* @param[in] src A reference to the source register tile.
*/
template<typename T2, typename U2, int _height, int _width, ducks::rt_layout::all layout>
__device__ static inline void copy(crt<T2, _height, _width, layout> &dst, const crt<U2, _height, _width, layout> &src) {
copy(dst.real, src.real);
copy(dst.imag, src.imag);
}
@@ -1,137 +0,0 @@
/**
* @file
* @brief Map operations between complex tiles.
*/
/**
* @brief Sets all elements of a complex tile to zero.
*
* @tparam T Complex tile type.
* @param dst[out] Destination tile where the result is stored.
*/
template<ducks::crt::all T>
__device__ static inline void zero(T &dst) {
zero(dst.real);
zero(dst.imag);
}
/**
* @brief Applies the exponential function to each element of a complex tile.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the exponential function on.
*/
template<ducks::crt::all T>
__device__ static inline void exp(T &dst, const T &src) {
using dtype = T::dtype;
dtype tmp;
// out of place storage
dtype rdst;
dtype idst;
// exp(a)
exp(rdst, src.real);
copy(idst, rdst);
// exp(a)cos(b) + exp(a)sin(b)i
cos(tmp, src.imag);
mul(rdst, rdst, tmp);
sin(tmp, src.imag);
mul(idst, idst, tmp);
copy(dst.real, rdst);
copy(dst.imag, idst);
}
/**
* @brief Adds two complex tiles element-wise.
*
* @tparam T Complex Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the addition.
* @param rhs[in] Right-hand side source tile for the addition.
*/
template<ducks::crt::all T>
__device__ static inline void add(T &dst, const T &lhs, const T &rhs) {
add(dst.real, lhs.real, rhs.real);
add(dst.imag, lhs.imag, rhs.imag);
}
/**
* @brief Subtracts two tiles element-wise.
*
* @tparam T Tile type.
* @tparam U Second operand type, which can be a tile or a scalar.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the subtraction.
* @param rhs[in] Right-hand side source tile for the subtraction.
*/
template<ducks::crt::all T>
__device__ static inline void sub(T &dst, const T &lhs, const T &rhs) {
sub(dst.real, lhs.real, rhs.real);
sub(dst.imag, lhs.imag, rhs.imag);
}
/**
* @brief Multiplies two tiles element-wise.
*
* @tparam T Complex tile type.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the multiplication.
* @param rhs[in] Right-hand side source tile for the multiplication.
*/
template<ducks::crt::all T>
__device__ static inline void mul(T &dst, const T &lhs, const T &rhs) {
using dtype = T::component;
dtype tmp;
// out of place storage regs
dtype rdst;
dtype idst;
// (a + bi) * (c + di) --> (ac - bd) + (ad + bc)i
// Real component
mul(rdst, lhs.real, rhs.real);
mul(tmp, lhs.imag, rhs.imag);
sub(rdst, rdst, tmp);
// Imag component
mul(idst, lhs.imag, rhs.real);
mul(tmp, lhs.real, rhs.imag);
add(idst, idst, tmp);
copy(dst.real, rdst);
copy(dst.imag, idst);
}
/**
* @brief Divides two tiles element-wise.
*
* @tparam T Complex tile type.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the division.
* @param rhs[in] Right-hand side source tile or scalar for the division.
*/
template<ducks::crt::all T>
__device__ static inline void div(T &dst, const T &lhs, const T &rhs) {
using dtype = T::dtype;
dtype tmp;
dtype denom;
// out of place storage regs
dtype rdst;
dtype idst;
// Calculate denom - square of b terms
mul(tmp, rhs.real, rhs.real);
mul(denom, rhs.imag, rhs.imag);
add(denom, tmp, denom);
// Real component
mul(rdst, lhs.real, rhs.real);
mul(tmp, lhs.imag, rhs.imag);
add(rdst, rdst, tmp);
// Imag component
mul(dst.imag, lhs.imag, rhs.real);
mul(tmp, lhs.real, rhs.imag);
sub(idst, idst, tmp);
// Divide components by denom
div(rdst, rdst, denom);
div(idst, idst, denom);
copy(dst.real, rdst);
copy(dst.imag, idst);
}
@@ -1,415 +0,0 @@
/**
* @file
* @brief Conversions between data layouts and types for register tiles.
*/
/* ---------- LAYOUT SWAPS ---------- */
/**
* @brief Perform a matrix transpose on a block of 8 bf16_2 elements using inline assembly.
*
* This low-level operation is utilized by higher-level layout swap functions to transpose
* the layout of bf16_2 elements within a register tile. The function leverages inline PTX
* assembly to efficiently swap the layout of the given block.
*
* @param[out] dst A reference to the destination bf16_2 element where the transposed result is stored.
* @param[in] src A reference to the source bf16_2 element to be transposed.
*/
__device__ static inline void swap_layout_8(bf16_2 &dst, const bf16_2 &src) {
KITTENS_CHECK_WARP
asm volatile (
"movmatrix.sync.aligned.m8n8.trans.b16 %0, %1;\n"
: "+r"(*(uint32_t*)(&dst))
: "r"(*(uint32_t*)(&src))
);
}
/**
* @brief Swaps the layout of a register base tile.
*
* This function swaps the layout of a register base tile by performing a series of layout swaps
* on its constituent bf16_2 elements. It is used to change the data layout within a register tile.
*
* @tparam T2 The data type of the register tile elements.
* @tparam layout The current layout of the register tile.
* @param dst[out] Reference to the destination register base tile where the result will be stored.
* @param src[in] Reference to the source register base tile to be swapped.
*/
template<typename T, ducks::rt_layout::all layout>
__device__ static inline void swap_layout(rt_base<T, typename ducks::rt_layout::transpose<layout>::type> &dst, const rt_base<T, layout> &src) {
swap_layout_8(dst.data[0], src.data[0]);
// technically this swap can be eliminated if we simply reinterpret the layout of the registers
// everywhere else in the code, but that feels... very likely to cause bugs and not worth it.
typename rt_base<T, layout>::T2 data1_cache = src.data[1]; // important for swap!
swap_layout_8(dst.data[1], src.data[2]);
swap_layout_8(dst.data[2], data1_cache);
swap_layout_8(dst.data[3], src.data[3]);
}
/**
* @brief Swaps the layout of a register tile.
*
* This function swaps the layout of a register tile by iterating over its height and width
* and performing layout swaps on each of its base elements.
*
* @tparam T2 The data type of the register tile elements.
* @tparam _height The height of the register tile.
* @tparam _width The width of the register tile.
* @tparam layout The current layout of the register tile.
* @param dst[out] Reference to the destination register tile where the result will be stored.
* @param src[in] Reference to the source register tile to be swapped.
*/
template<typename T2, int _height, int _width, ducks::rt_layout::all layout>
__device__ static inline void swap_layout(rt<T2, _height, _width, typename ducks::rt_layout::transpose<layout>::type> &dst, const rt<T2, _height, _width, layout> &src) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
swap_layout(dst.tiles[i][j], src.tiles[i][j]);
}
}
}
/**
* @brief Swaps the layout of a register base tile in place.
*
* This function swaps the layout of a register base tile in place by casting it to the
* transposed layout type and then performing the layout swap.
*
* @tparam T2 The data type of the register tile elements.
* @tparam layout The current layout of the register tile.
* @param src[in] Reference to the register base tile to be swapped in place.
* @return A reference to the swapped register base tile.
*/
template<typename T2, ducks::rt_layout::all layout>
__device__ static inline rt_base<T2, typename ducks::rt_layout::transpose<layout>::type>& swap_layout_inplace(const rt_base<T2, layout> &src) {
rt_base<T2, typename ducks::rt_layout::transpose<layout>::type> &dst = *(rt_base<T2, typename ducks::rt_layout::transpose<layout>::type>*)(&src);
swap_layout(dst, src);
return dst;
}
/**
* @brief Swaps the layout of a register tile in place.
*
* This function swaps the layout of a register tile in place by iterating over its height and width
* and performing in-place layout swaps on each of its base elements.
*
* @tparam T2 The data type of the register tile elements.
* @tparam _height The height of the register tile.
* @tparam _width The width of the register tile.
* @tparam layout The current layout of the register tile.
* @param tile[in,out] Reference to the register tile to be swapped in place.
* @return A reference to the swapped register tile.
*/
template<typename T2, int _rows, int _cols, ducks::rt_layout::all layout>
__device__ static inline rt<T2, _rows, _cols, typename ducks::rt_layout::transpose<layout>::type>& swap_layout_inplace(rt<T2, _rows, _cols, layout> &tile) {
#pragma unroll
for(int i = 0; i < tile.height; i++) {
#pragma unroll
for(int j = 0; j < tile.width; j++) {
swap_layout_inplace(tile.tiles[i][j]);
}
}
return *(rt<T2, _rows, _cols, typename ducks::rt_layout::transpose<layout>::type>*)(&tile);
}
/* ---------- TRANSPOSE ---------- */
/**
* @brief Transposes a register base tile.
*
* @tparam T2 The data type of the register tile elements.
* @tparam layout The current layout of the register tile.
* @param dst[out] Reference to the register tile in which to store the transposed src.
* @param src[in] Reference to the register base tile to be transposed.
*/
template<typename T, ducks::rt_layout::all layout>
__device__ static inline void transpose(rt_base<T, layout> &dst, const rt_base<T, layout> &src) {
swap_layout_8(dst.data[0], src.data[0]);
// technically this swap can be eliminated if we simply reinterpret the layout of the registers
// everywhere else in the code, but that feels... very likely to cause bugs and not worth it.
typename rt_base<T, layout>::T2 data1_cache = src.data[1]; // important for swap!
swap_layout_8(dst.data[1], src.data[2]);
swap_layout_8(dst.data[2], data1_cache);
swap_layout_8(dst.data[3], src.data[3]);
}
/**
* @brief Transposes a register tile.
*
* This function is marked "sep", which means that the registers underlying dst MUST be separate
* from the registers underlying src.
*
* @tparam T2 The data type of the register tile elements.
* @tparam _height The height of the src register tile, and the width of the dst tile.
* @tparam _width The width of the src register tile, and the height of the dst tile.
* @tparam layout The layout of the register tile.
* @param dst[out] Reference to the register tile in which to store the transposed src.
* @param src[in] Reference to the register tile to be transposed.
*/
template<ducks::rt::all RT>
__device__ static inline void transpose_sep(RT &dst, const rt<typename RT::T, RT::cols, RT::rows, typename RT::layout> &src) {
#pragma unroll
for(int i = 0; i < RT::height; i++) {
#pragma unroll
for(int j = 0; j < RT::width; j++) {
transpose(dst.tiles[i][j], src.tiles[j][i]);
}
}
}
/**
* @brief Transposes a register base tile in-place.
*
* @tparam T2 The data type of the register base tile elements.
* @tparam layout The current layout of the register base tile.
* @param src[in] Reference to the register tile to be transposed.
* @return A reference to the transposed register base tile.
*/
template<typename T2, ducks::rt_layout::all layout>
__device__ static inline rt_base<T2, layout>& transpose_inplace(rt_base<T2, layout> &src) {
transpose(src, src);
return src;
}
/**
* @brief Transposes a square register tile in-place.
*
* @tparam T2 The data type of the register tile elements.
* @tparam _height The height (in units of 16) of the src register tile, and the width of the dst tile. (Must be the same as _width.)
* @tparam _width The width (in units of 16) of the src register tile, and the height of the dst tile. (Must be the same as _height.)
* @tparam layout The current layout of the register tile.
* @param src[in] Reference to the register tile to be transposed.
* @return A reference to the transposed register tile.
*/
template<typename T2, int _rows, int _cols, ducks::rt_layout::all layout>
__device__ static inline rt<T2, _rows, _cols, layout>& transpose_inplace(rt<T2, _rows, _cols, layout> &tile) {
static_assert(_cols == _rows, "in-place register tile transpose is only allowed for square tiles.");
#pragma unroll
for(int i = 0; i < tile.height; i++) {
#pragma unroll
for(int j = 0; j < i; j++) {
rt_base<T2, layout> tmp;
copy(tmp, tile.tiles[i][j]);
transpose(tile.tiles[i][j], tile.tiles[j][i]);
transpose(tile.tiles[j][i], tmp);
}
transpose_inplace(tile.tiles[i][i]);
}
return tile;
}
/* ---------- TYPE SWAPS ---------- */
/**
* @brief Copies a register base tile, converting the underlying type if necessary.
*
* @tparam T2 The data type of the destination register elements.
* @tparam U2 The data type of the source register elements.
* @tparam layout The current layout of the register base tile.
* @param[out] dst A reference to the destination register base tile.
* @param[in] src A reference to the source register base tile.
*/
template<typename T, typename U, ducks::rt_layout::all layout>
__device__ static inline void copy(rt_base<T, layout> &dst, const rt_base<U, layout> &src) {
using T2 = typename base_types::packing<T>::packed_type;
using U2 = typename base_types::packing<U>::packed_type;
#pragma unroll
for(int k = 0; k < dst.packed_per_thread; k++) {
dst.data[k] = base_types::convertor<T2, U2>::convert(src.data[k]);
}
}
#ifdef KITTENS_HOPPER
/**
* @brief Copies a register tile, converting the underlying type if necessary.
*
* @tparam T2 The data type of the destination register elements.
* @tparam U2 The data type of the source register elements.
* @tparam _height The height (in units of 16) of the register tiles.
* @tparam _width The width (in units of 16) of the register tiles.
* @tparam layout The current layout of the register tile.
* @param[out] dst A reference to the destination register tile.
* @param[in] src A reference to the source register tile.
*/
template<typename T2, typename U2, int _height, int _width, ducks::rt_layout::all layout>
__device__ static inline void copy(rt<T2, _height, _width, layout> &dst, const rt<U2, _height, _width, layout> &src) {
if constexpr (
(std::is_same_v<U2, float> && std::is_same_v<T2, fp8e4m3>) ||
(std::is_same_v<U2, float> && std::is_same_v<T2, fp8e5m2>) ||
(std::is_same_v<U2, kittens::bf16> && std::is_same_v<T2, fp8e4m3>) ||
(std::is_same_v<U2, kittens::bf16> && std::is_same_v<T2, fp8e5m2>) ||
(std::is_same_v<U2, half> && std::is_same_v<T2, fp8e4m3>) ||
(std::is_same_v<U2, half> && std::is_same_v<T2, fp8e5m2>)
) {
// FLOAT (SRC -- 1H x 2W) to FP8 (DST -- 1H x 1W)
int laneid = threadIdx.x % 32;
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.tiles[0][0].packed_per_thread; k++) {
// check for half, float, bf16
using src_t = std::conditional_t<std::is_same_v<U2, float>, float2, std::conditional_t<std::is_same_v<U2, kittens::bf16>, bf16_2, half2>>;
src_t val1, val2;
// Put something up for adoption
if (laneid % 2 == 0) {
// put up src left core matrix first as 0, 2
val1 = src.tiles[i][2*j + k/2].data[(k%2)+0];
val2 = src.tiles[i][2*j + k/2].data[(k%2)+2];
} else {
// put up src right core matrix first as 1, 3
val1 = src.tiles[i][2*j + k/2].data[(k%2)+2];
val2 = src.tiles[i][2*j + k/2].data[(k%2)+0];
}
// Shuffle first 4 floats
int row_mask = 4 * ( laneid / 4 );
int row_offset = row_mask + ( (laneid-row_mask) / 2 ) + ( laneid % 2 );
int src_offset = (laneid % 2 == 0 ) ? row_offset + 0 : ( row_offset + 1 );
src_t val01 = packed_shfl_sync(MASK_ALL, val1, src_offset); // Get from even thread
int src_offset2 = (laneid % 4 < 2 ) ? src_offset + 1 : (src_offset - 1);
src_t val23 = packed_shfl_sync(MASK_ALL, val2, src_offset2); // Get from odd thread
// Convert to fp8e4m3_4
float4 f4;
using fp8_4_t = std::conditional_t<std::is_same_v<T2, fp8e4m3>, fp8e4m3_4, fp8e5m2_4>;
fp8_4_t f4_fp8;
if ( laneid % 4 < 2 ) {
f4.x = val01.x; // Thread 2N's first value
f4.y = val01.y; // Thread 2N's second value
f4.z = val23.x; // Thread 2N+1's first value
f4.w = val23.y; // Thread 2N+1's second value
f4_fp8 = base_types::convertor<fp8_4_t, float4>::convert(f4);
dst.tiles[i][j].data[k] = f4_fp8;
} else {
f4.x = val23.x; // Thread 2N+1's first value
f4.y = val23.y; // Thread 2N+1's second value
f4.z = val01.x; // Thread 2N's first value
f4.w = val01.y; // Thread 2N's second value
f4_fp8 = base_types::convertor<fp8_4_t, float4>::convert(f4);
dst.tiles[i][j].data[k] = f4_fp8;
}
}
}
}
}
else if constexpr (
(std::is_same_v<U2, fp8e4m3> && std::is_same_v<T2, float>) ||
(std::is_same_v<U2, fp8e5m2> && std::is_same_v<T2, float>) ||
(std::is_same_v<U2, fp8e4m3> && std::is_same_v<T2, kittens::bf16>) ||
(std::is_same_v<U2, fp8e5m2> && std::is_same_v<T2, kittens::bf16>) ||
(std::is_same_v<U2, fp8e4m3> && std::is_same_v<T2, half>) ||
(std::is_same_v<U2, fp8e5m2> && std::is_same_v<T2, half>)
) {
// FP8 (SRC -- 1H x 1W) to FLOAT (DST -- 1H x 2W)
int laneid = threadIdx.x % 32;
#pragma unroll
for(int i = 0; i < src.height; i++) {
#pragma unroll
for(int j = 0; j < src.width; j++) {
#pragma unroll
for(int k = 0; k < src.tiles[0][0].packed_per_thread; k++) {
int dst_j = 2*j + k/2;
// Put something up for adoption
using fp8_4_t = std::conditional_t<std::is_same_v<U2, fp8e4m3>, fp8e4m3_4, fp8e5m2_4>;
fp8_4_t val = src.tiles[i][j].data[k];
float4 f4 = base_types::convertor<float4, fp8_4_t>::convert(val);
float2 f2_0, f2_1;
if ( laneid % 4 < 2 ) { // src 0 and 1 should put up .x and .y first
f2_0 = make_float2(f4.x, f4.y);
f2_1 = make_float2(f4.z, f4.w);
}
else { // src 2 and 3 should put up .z and .w first
f2_0 = make_float2(f4.z, f4.w);
f2_1 = make_float2(f4.x, f4.y);
}
int row_offset = 4 * (laneid/4) + (laneid%2) * 2 + (laneid%4) / 2;
float2 f2_0_shfl = packed_shfl_sync(MASK_ALL, f2_0, row_offset);
float2 f2_1_shfl = packed_shfl_sync(MASK_ALL, f2_1, row_offset^2);
// convert to dst type if needed
using dst_t = std::conditional_t<std::is_same_v<T2, float>, float2, std::conditional_t<std::is_same_v<T2, kittens::bf16>, bf16_2, half2>>;
if constexpr (!(std::is_same_v<T2, float>)) {
dst_t f2_0_shfl_t = base_types::convertor<dst_t, float2>::convert(f2_0_shfl);
dst_t f2_1_shfl_t = base_types::convertor<dst_t, float2>::convert(f2_1_shfl);
if (laneid % 2 == 0) {
dst.tiles[i][dst_j].data[(k%2)+0] = f2_0_shfl_t;
dst.tiles[i][dst_j].data[(k%2)+2] = f2_1_shfl_t;
} else {
dst.tiles[i][dst_j].data[(k%2)+0] = f2_1_shfl_t;
dst.tiles[i][dst_j].data[(k%2)+2] = f2_0_shfl_t;
}
} else {
if (laneid % 2 == 0) {
dst.tiles[i][dst_j].data[(k%2)+0] = f2_0_shfl;
dst.tiles[i][dst_j].data[(k%2)+2] = f2_1_shfl;
} else {
dst.tiles[i][dst_j].data[(k%2)+0] = f2_1_shfl;
dst.tiles[i][dst_j].data[(k%2)+2] = f2_0_shfl;
}
}
}
}
}
}
// default case where the layouts map 1:1 in thread ownership logic
else {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
copy(dst.tiles[i][j], src.tiles[i][j]);
}
}
}
}
#else
/**
* @brief Copies a register tile, converting the underlying type if necessary.
*
* @tparam T2 The data type of the destination register elements.
* @tparam U2 The data type of the source register elements.
* @tparam _height The height (in units of 16) of the register tiles.
* @tparam _width The width (in units of 16) of the register tiles.
* @tparam layout The current layout of the register tile.
* @param[out] dst A reference to the destination register tile.
* @param[in] src A reference to the source register tile.
*/
template<typename T2, typename U2, int _height, int _width, ducks::rt_layout::all layout>
__device__ static inline void copy(rt<T2, _height, _width, layout> &dst, const rt<U2, _height, _width, layout> &src) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
copy(dst.tiles[i][j], src.tiles[i][j]);
}
}
}
#endif
/* ---------- SUBTILE ---------- */
/**
* @brief Returns a reference to a subtile of the given tile.
*
* @tparam subtile_height The height of the subtile.
* @tparam RT The type of the input tile, which must satisfy the ducks::rt::all concept.
* @param src The input tile.
* @param idx The coord of the subtile.
* @return A reference to the subtile.
*
* @note The subtile height must evenly divide the tile height.
*/
template<int subtile_rows, ducks::rt::all RT>
__device__ static inline rt<typename RT::T, subtile_rows, RT::cols, typename RT::layout> &subtile_inplace(RT & src, int idx) {
KITTENS_CHECK_WARP
using T = typename RT::T;
static_assert(RT::height % (subtile_rows / TILE_ROW_DIM<T>) == 0, "subtile height should evenly divide tile height.");
return reinterpret_cast<rt<typename RT::T, subtile_rows, RT::cols, typename RT::layout>&>(
src.tiles[idx*(subtile_rows / TILE_ROW_DIM<T>)]
);
}
@@ -1,836 +0,0 @@
/**
* @file
* @brief Map operations: between tiles, and those which apply vectors to tiles.
*/
/* ---------- Uniform tile maps (independent of layout) ---------- */
/**
* @brief Applies a unary operation to each element of a tile.
*
* @tparam op Unary operation to apply.
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the operation on.
*/
template<typename op, ducks::rt::all T>
__device__ static inline void unary_map(T &dst, const T &src) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k++) {
dst.tiles[i][j].data[k] = op::template op<typename T::dtype>(src.tiles[i][j].data[k]);
}
}
}
}
/**
* @brief Applies a binary operation to each element of a tile with a scalar parameter.
*
* @tparam op Binary operation to apply.
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the operation on.
* @param param[in] Scalar parameter for the binary operation.
*/
template<typename op, ducks::rt::all T>
__device__ static inline void bin_map(T &dst, const T &src, const typename T::dtype &param) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k++) {
dst.tiles[i][j].data[k] = op::template op<typename T::dtype>(src.tiles[i][j].data[k], param);
}
}
}
}
/**
* @brief Applies a binary operation to each element of a tile with an unpacked scalar parameter.
*
* @tparam op Binary operation to apply.
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the operation on.
* @param param[in] Unpacked scalar parameter for the binary operation.
*/
template<typename op, ducks::rt::all T>
__device__ static inline void bin_map(T &dst, const T &src, const typename base_types::packing<typename T::dtype>::unpacked_type &param) {
// The optimizing compiler should eliminate this pack in the 32-bit case but not in the 16-bit case
bin_map<op, T>(dst, src, base_types::packing<typename T::dtype>::pack(param));
}
/**
* @brief Applies a binary operation element-wise between two tiles.
*
* @tparam op Binary operation to apply.
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the operation.
* @param rhs[in] Right-hand side source tile for the operation.
*/
template<typename op, ducks::rt::all T>
__device__ static inline void bin_map(T &dst, const T &lhs, const T &rhs) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k++) {
dst.tiles[i][j].data[k] = op::template op<typename T::dtype>(lhs.tiles[i][j].data[k], rhs.tiles[i][j].data[k]);
}
}
}
}
template<ducks::rt::all RT, typename Lambda>
__device__ static inline void apply(RT &dst, const RT &src, Lambda &&lambda) {
int row_offset = 0;
if constexpr(GROUP_WARPS > 1) {
row_offset = warpid()*RT::height;
}
static_assert(sizeof(RT::T) != 1, "Cannot apply lambda to 8-bit types");
if constexpr (ducks::rt::row_layout<RT>) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k++) {
int row = row_offset + i*TILE_ROW_DIM<typename RT::T> + (k%2) * (TILE_ROW_DIM<typename RT::T>/2) + ::kittens::laneid()/4;
int col = j*TILE_COL_DIM<typename RT::T> + (k/2) * (TILE_COL_DIM<typename RT::T>/2) + (::kittens::laneid()%4)*2;
dst.tiles[i][j].data[k].x = lambda(row, col+0, src.tiles[i][j].data[k].x);
dst.tiles[i][j].data[k].y = lambda(row, col+1, src.tiles[i][j].data[k].y);
}
}
}
}
else {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k++) {
int row = row_offset + i*TILE_ROW_DIM<typename RT::T> + (k/2) * (TILE_ROW_DIM<typename RT::T>/2) + (::kittens::laneid()%4)*2;
int col = j*TILE_COL_DIM<typename RT::T> + (k%2) * (TILE_COL_DIM<typename RT::T>/2) + ::kittens::laneid()/4;
dst.tiles[i][j].data[k].x = lambda(row+0, col, src.tiles[i][j].data[k].x);
dst.tiles[i][j].data[k].y = lambda(row+1, col, src.tiles[i][j].data[k].y);
}
}
}
}
}
template<ducks::rt::all RT, typename Lambda>
__device__ static inline RT apply(const RT &src, Lambda &&lambda) {
RT dst;
apply<RT, Lambda>(dst, src, std::forward<Lambda>(lambda));
return dst;
}
/* ---------- Row tile maps ----------*/
/**
* @brief Applies an operation across the rows of a tile in a row-major layout.
*
* @tparam op Operation to apply.
* @tparam T Tile type with row-major layout.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the operation on.
* @param row_values[in] Column vector containing values to apply across each row.
*/
template<typename op, ducks::rt::row_layout T, ducks::rv::all V>
__device__ static inline void row_map(T &dst, const T &src, const V &row_values) {
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::col_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::height); // compatible size
using dtype = T::dtype;
#pragma unroll
for(int i = 0; i < dst.height; i++) {
dtype packed_top_row = base_types::packing<dtype>::pack(row_values[i][0].x); // first value in eager mode
dtype packed_bottom_row = base_types::packing<dtype>::pack(row_values[i][0].y); // second value in eager mode
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k+=2) {
dst.tiles[i][j].data[k+0] = op::template op<dtype>(src.tiles[i][j].data[k+0], packed_top_row);
dst.tiles[i][j].data[k+1] = op::template op<dtype>(src.tiles[i][j].data[k+1], packed_bottom_row);
}
}
}
}
/**
* @brief Applies an operation across the rows of a tile in a column-major layout.
*
* @tparam op Operation to apply.
* @tparam T Tile type with column-major layout.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the operation on.
* @param row_values[in] Column vector containing values to apply across each row.
*/
template<typename op, ducks::rt::col_layout T, ducks::rv::all V>
__device__ static inline void row_map(T &dst, const T &src, const V &row_values) {
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::col_vec_layout>); // compatible layout
static_assert(V::outer_dim == T::height); // compatible size
using dtype = T::dtype;
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile/2; k++) {
dst.tiles[i][j].data[k+0] = op::template op<dtype>(src.tiles[i][j].data[k+0], row_values[i][0]);
dst.tiles[i][j].data[k+2] = op::template op<dtype>(src.tiles[i][j].data[k+2], row_values[i][1]);
}
}
}
}
// Three-operand row map. Mostly useful for FMA instructions.
/**
* @brief Applies an operation across the rows of two tiles in a row-major layout, using a third operand.
*
* @tparam op Operation to apply.
* @tparam T Tile type with row-major layout.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param a[in] First source tile to apply the operation on.
* @param b[in] Second source tile to apply the operation on.
* @param row_values[in] Column vector containing values to apply across each row.
*/
template<typename op, ducks::rt::row_layout T, ducks::rv::all V>
__device__ static inline void row_map(T &dst, const T &a, const T &b, const V &row_values) {
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::col_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::height); // compatible size
using dtype = T::dtype;
#pragma unroll
for(int i = 0; i < dst.height; i++) {
dtype packed_top_row = base_types::packing<dtype>::pack(row_values[i][0].x); // first value in eager mode
dtype packed_bottom_row = base_types::packing<dtype>::pack(row_values[i][0].y); // second value in eager mode
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k+=2) {
dst.tiles[i][j].data[k+0] = op::template op<dtype>(a.tiles[i][j].data[k+0], b.tiles[i][j].data[k+0], packed_top_row);
dst.tiles[i][j].data[k+1] = op::template op<dtype>(a.tiles[i][j].data[k+1], b.tiles[i][j].data[k+1], packed_bottom_row);
}
}
}
}
/**
* @brief Applies an operation across the rows of two tiles in a column-major layout, using a third operand.
*
* @tparam op Operation to apply.
* @tparam T Tile type with column-major layout.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param a[in] First source tile to apply the operation on.
* @param b[in] Second source tile to apply the operation on.
* @param row_values[in] Column vector containing values to apply across each row.
*/
template<typename op, ducks::rt::col_layout T, ducks::rv::all V>
__device__ static inline void row_map(T &dst, const T &a, const T &b, const V &row_values) {
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::col_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::height); // compatible size
using dtype = T::dtype;
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile/2; k++) {
dst.tiles[i][j].data[k+0] = op::template op<dtype>(a.tiles[i][j].data[k+0], b.tiles[i][j].data[k+0], row_values[i][0]);
dst.tiles[i][j].data[k+2] = op::template op<dtype>(a.tiles[i][j].data[k+2], b.tiles[i][j].data[k+2], row_values[i][1]);
}
}
}
}
/* ---------- Col major tile maps ----------*/
/**
* @brief Applies an operation across the columns of a tile in a row-major layout.
*
* @tparam op Operation to apply.
* @tparam T Tile type with row-major layout.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the operation on.
* @param col_values[in] Row vector containing values to apply across each column.
*/
template<typename op, ducks::rt::row_layout T, ducks::rv::all V>
__device__ static inline void col_map(T &dst, const T &src, const V &col_values) {
KITTENS_CHECK_WARP
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::row_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::width); // compatible size
using dtype = T::dtype;
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile/2; k++) {
dst.tiles[i][j].data[k+0] = op::template op<dtype>(src.tiles[i][j].data[k+0], col_values[j][0]);
dst.tiles[i][j].data[k+2] = op::template op<dtype>(src.tiles[i][j].data[k+2], col_values[j][1]);
}
}
}
}
/**
* @brief Applies an operation across the columns of a tile in a column-major layout.
*
* @tparam op Operation to apply.
* @tparam T Tile type with column-major layout.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the operation on.
* @param col_values[in] Row vector containing values to apply across each column.
*/
template<typename op, ducks::rt::col_layout T, ducks::rv::all V>
__device__ static inline void col_map(T &dst, const T &src, const V &col_values) {
KITTENS_CHECK_WARP
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::row_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::width); // compatible size
using dtype = T::dtype;
#pragma unroll
for(int j = 0; j < dst.width; j++) {
dtype packed_left_col = base_types::packing<dtype>::pack(col_values[j][0].x); // first value in eager mode
dtype packed_right_col = base_types::packing<dtype>::pack(col_values[j][0].y); // second value in eager mode
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k+=2) {
dst.tiles[i][j].data[k+0] = op::template op<dtype>(src.tiles[i][j].data[k+0], packed_left_col);
dst.tiles[i][j].data[k+1] = op::template op<dtype>(src.tiles[i][j].data[k+1], packed_right_col);
}
}
}
}
// Three-operand col map
/**
* @brief Applies an operation across the columns of two tiles in a row-major layout, using a third operand.
*
* @tparam op Operation to apply.
* @tparam T Tile type with row-major layout.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param a[in] First source tile to apply the operation on.
* @param b[in] Second source tile to apply the operation on.
* @param col_values[in] Row vector containing values to apply across each column.
*/
template<typename op, ducks::rt::row_layout T, ducks::rv::all V>
__device__ static inline void col_map(T &dst, const T &a, const T &b, const V &col_values) {
KITTENS_CHECK_WARP
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::row_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::width); // compatible size
using dtype = T::dtype;
#pragma unroll
for(int j = 0; j < dst.width; j++) {
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile/2; k++) {
dst.tiles[i][j].data[k+0] = op::template op<dtype>(a.tiles[i][j].data[k+0], b.tiles[i][j].data[k+0], col_values[j][0]);
dst.tiles[i][j].data[k+2] = op::template op<dtype>(a.tiles[i][j].data[k+2], b.tiles[i][j].data[k+2], col_values[j][1]);
}
}
}
}
/**
* @brief Applies an operation across the columns of two tiles in a column-major layout, using a third operand.
*
* @tparam op Operation to apply.
* @tparam T Tile type with column-major layout.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param a[in] First source tile to apply the operation on.
* @param b[in] Second source tile to apply the operation on.
* @param col_values[in] Row vector containing values to apply across each column.
*/
template<typename op, ducks::rt::col_layout T, ducks::rv::all V>
__device__ static inline void col_map(T &dst, const T &a, const T &b, const V &col_values) {
KITTENS_CHECK_WARP
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::row_vec_layout>); // compatible layout
static_assert(V::outer_dim == T::width); // compatible size
using dtype = T::dtype;
#pragma unroll
for(int j = 0; j < dst.width; j++) {
dtype packed_left_col = base_types::packing<dtype>::pack(col_values[j][0].x); // first value in eager mode
dtype packed_right_col = base_types::packing<dtype>::pack(col_values[j][0].y); // second value in eager mode
#pragma unroll
for(int i = 0; i < dst.height; i++) {
#pragma unroll
for(int k = 0; k < dst.packed_per_tile; k+=2) {
dst.tiles[i][j].data[k+0] = op::template op<dtype>(a.tiles[i][j].data[k+0], b.tiles[i][j].data[k+0], packed_left_col);
dst.tiles[i][j].data[k+1] = op::template op<dtype>(a.tiles[i][j].data[k+1], b.tiles[i][j].data[k+1], packed_right_col);
}
}
}
}
/* ---------- WRAPPERS FOR PRETTINESS ---------- */
// All of the annoying qualifiers *should* be automatically inferred during compile-time.
// So, syntax should just be kittens::add_row(tile, colvec);
/**
* @brief Sets all elements of a tile to zero.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
*/
template<ducks::rt::all T>
__device__ static inline void zero(T &dst) {
unary_map<base_ops::zero, T>(dst, dst);
}
/**
* @brief Sets all elements of a tile to one.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
*/
template<ducks::rt::all T>
__device__ static inline void one(T &dst) {
unary_map<base_ops::one, T>(dst, dst);
}
/**
* @brief Sets all elements of a tile to positive infinity.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
*/
template<ducks::rt::all T>
__device__ static inline void pos_infty(T &dst) {
unary_map<base_ops::pos_infty, T>(dst, dst);
}
/**
* @brief Sets all elements of a tile to negative infinity.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
*/
template<ducks::rt::all T>
__device__ static inline void neg_infty(T &dst) {
unary_map<base_ops::neg_infty, T>(dst, dst);
}
/**
* @brief Applies the exponential function to each element of a tile.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the exponential function on.
*/
template<ducks::rt::all T>
__device__ static inline void exp(T &dst, const T &src) {
unary_map<base_ops::exp, T>(dst, src);
}
template<ducks::rt::all T>
__device__ static inline T exp(const T &src) {
T dst;
exp(dst, src);
return dst;
}
/**
* @brief Applies the exponential function to each element of a tile, in base 2.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the exponential function on.
*/
template<ducks::rt::all T>
__device__ static inline void exp2(T &dst, const T &src) {
unary_map<base_ops::exp2, T>(dst, src);
}
template<ducks::rt::all T>
__device__ static inline T exp2(const T &src) {
T dst;
exp2(dst, src);
return dst;
}
/**
* @brief Applies the natural logarithm function to each element of a tile.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the natural logarithm function on.
*/
template<ducks::rt::all T>
__device__ static inline void log(T &dst, const T &src) {
unary_map<base_ops::log, T>(dst, src);
}
template<ducks::rt::all T>
__device__ static inline T log(const T &src) {
T dst;
log(dst, src);
return dst;
}
/**
* @brief Applies the logarithm base 2 function to each element of a tile.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the logarithm base 2 function on.
*/
template<ducks::rt::all T>
__device__ static inline void log2(T &dst, const T &src) {
unary_map<base_ops::log2, T>(dst, src);
}
template<ducks::rt::all T>
__device__ static inline T log2(const T &src) {
T dst;
log2(dst, src);
return dst;
}
/**
* @brief Applies the absolute value function to each element of a tile.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the absolute value function on.
*/
template<ducks::rt::all T>
__device__ static inline void abs(T &dst, const T &src) {
unary_map<base_ops::abs, T>(dst, src);
}
template<ducks::rt::all T>
__device__ static inline T abs(const T &src) {
T dst;
abs(dst, src);
return dst;
}
/**
* @brief Applies the rectified linear unit (ReLU) function to each element of a tile.
*
* @tparam T Tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the ReLU function on.
*/
template<ducks::rt::all T>
__device__ static inline void relu(T &dst, const T &src) {
unary_map<base_ops::relu, T>(dst, src);
}
template<ducks::rt::all T>
__device__ static inline T relu(const T &src) {
T dst;
relu(dst, src);
return dst;
}
/**
* @brief Copies the elements from one tile to another.
*
* @tparam T Destination tile type.
* @tparam U Source tile type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to copy from.
*/
template<ducks::rt::all T, typename U>
__device__ static inline void copy(T &dst, const U &src) {
bin_map<base_ops::copy2, T>(dst, src);
}
/**
* @brief Applies the max operation element-wise between two tiles or a tile and a scalar.
*
* @tparam T Tile type.
* @tparam U Second operand type, which can be a tile or a scalar.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the operation.
* @param rhs[in] Right-hand side source tile or scalar for the operation.
*/
template<ducks::rt::all T, typename U>
__device__ static inline void max(T &dst, const T &lhs, const U &rhs) {
bin_map<base_ops::max, T>(dst, lhs, rhs);
}
template<ducks::rt::all T, typename U>
__device__ static inline T max(const T &lhs, const U &rhs) {
T dst;
max(dst, lhs, rhs);
return dst;
}
/**
* @brief Applies the min operation element-wise between two tiles or a tile and a scalar.
*
* @tparam T Tile type.
* @tparam U Second operand type, which can be a tile or a scalar.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the operation.
* @param rhs[in] Right-hand side source tile or scalar for the operation.
*/
template<ducks::rt::all T, typename U>
__device__ static inline void min(T &dst, const T &lhs, const U &rhs) {
bin_map<base_ops::min, T>(dst, lhs, rhs);
}
template<ducks::rt::all T, typename U>
__device__ static inline T min(const T &lhs, const U &rhs) {
T dst;
min(dst, lhs, rhs);
return dst;
}
/**
* @brief Adds two tiles element-wise or adds a scalar to each element of a tile.
*
* @tparam T Tile type.
* @tparam U Second operand type, which can be a tile or a scalar.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the addition.
* @param rhs[in] Right-hand side source tile or scalar for the addition.
*/
template<ducks::rt::all T, typename U>
__device__ static inline void add(T &dst, const T &lhs, const U &rhs) {
bin_map<base_ops::sum, T>(dst, lhs, rhs);
}
/**
* @brief Subtracts two tiles element-wise or subtracts a scalar from each element of a tile.
*
* @tparam T Tile type.
* @tparam U Second operand type, which can be a tile or a scalar.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the subtraction.
* @param rhs[in] Right-hand side source tile or scalar for the subtraction.
*/
template<ducks::rt::all T, typename U>
__device__ static inline void sub(T &dst, const T &lhs, const U &rhs) {
bin_map<base_ops::sub, T>(dst, lhs, rhs);
}
/**
* @brief Multiplies two tiles element-wise or multiplies each element of a tile by a scalar.
*
* @tparam T Tile type.
* @tparam U Second operand type, which can be a tile or a scalar.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the multiplication.
* @param rhs[in] Right-hand side source tile or scalar for the multiplication.
*/
template<ducks::rt::all T, typename U>
__device__ static inline void mul(T &dst, const T &lhs, const U &rhs) {
bin_map<base_ops::mul, T>(dst, lhs, rhs);
}
/**
* @brief Divides two tiles element-wise or divides each element of a tile by a scalar.
*
* @tparam T Tile type.
* @tparam U Second operand type, which can be a tile or a scalar.
* @param dst[out] Destination tile where the result is stored.
* @param lhs[in] Left-hand side source tile for the division.
* @param rhs[in] Right-hand side source tile or scalar for the division.
*/
template<ducks::rt::all T, typename U>
__device__ static inline void div(T &dst, const T &lhs, const U &rhs) {
bin_map<base_ops::div, T>(dst, lhs, rhs);
}
/**
* @brief Adds row values to each row of a tile.
*
* @tparam T Tile type.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the addition on.
* @param row_values[in] Column vector containing values to add to each row.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void add_row(T &dst, const T &src, const V &row_values) {
row_map<base_ops::sum, T, V>(dst, src, row_values);
}
/**
* @brief Subtracts row values from each row of a tile.
*
* @tparam T Tile type.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the subtraction on.
* @param row_values[in] Column vector containing values to subtract from each row.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void sub_row(T &dst, const T &src, const V &row_values) {
row_map<base_ops::sub, T, V>(dst, src, row_values);
}
/**
* @brief Multiplies each row of a tile by row values.
*
* @tparam T Tile type.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the multiplication on.
* @param row_values[in] Column vector containing values to multiply each row by.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void mul_row(T &dst, const T &src, const V &row_values) {
row_map<base_ops::mul, T, V>(dst, src, row_values);
}
/**
* @brief Divides each row of a tile by row values.
*
* @tparam T Tile type.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the division on.
* @param row_values[in] Column vector containing values to divide each row by.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void div_row(T &dst, const T &src, const V &row_values) {
row_map<base_ops::div, T, V>(dst, src, row_values);
}
/**
* @brief Broadcast a vector into into a tile's rows.
*
* @tparam T Tile type.
* @tparam V Column vector type.
* @param dst[out] Destination tile where the result is stored.
* @param row_values[in] Column vector containing values to broadcast into rows.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void broadcast_row(T &dst, const V &row_values) {
row_map<base_ops::copy2, T, V>(dst, dst, row_values);
}
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline T broadcast_row(const V &row_values) {
T dst;
broadcast_row(dst, row_values);
return dst;
}
// col maps
/**
* @brief Adds column values to each column of a tile.
*
* @tparam T Tile type.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the addition on.
* @param col_values[in] Row vector containing values to add to each column.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void add_col(T &dst, const T &src, const V &col_values) {
col_map<base_ops::sum, T, V>(dst, src, col_values);
}
/**
* @brief Subtracts column values from each column of a tile.
*
* @tparam T Tile type.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the subtraction on.
* @param col_values[in] Row vector containing values to subtract from each column.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void sub_col(T &dst, const T &src, const V &col_values) {
col_map<base_ops::sub, T, V>(dst, src, col_values);
}
/**
* @brief Multiplies each column of a tile by column values.
*
* @tparam T Tile type.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the multiplication on.
* @param col_values[in] Row vector containing values to multiply each column by.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void mul_col(T &dst, const T &src, const V &col_values) {
col_map<base_ops::mul, T, V>(dst, src, col_values);
}
/**
* @brief Divides each column of a tile by column values.
*
* @tparam T Tile type.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param src[in] Source tile to apply the division on.
* @param col_values[in] Row vector containing values to divide each column by.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void div_col(T &dst, const T &src, const V &col_values) {
col_map<base_ops::div, T, V>(dst, src, col_values);
}
/**
* @brief Broadcast a vector into into a tile's columns.
*
* @tparam T Tile type.
* @tparam V Row vector type.
* @param dst[out] Destination tile where the result is stored.
* @param row_values[in] Row vector containing values to broadcast into cols.
*/
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline void broadcast_col(T &dst, const V &col_values) {
col_map<base_ops::copy2, T, V>(dst, dst, col_values);
}
template<ducks::rt::all T, ducks::rv::all V>
__device__ static inline T broadcast_col(const V &col_values) {
T dst;
broadcast_col(dst, col_values);
return dst;
}
// Triangular masks
template<ducks::rt::all RT>
__device__ static inline void tril(RT &dst, const RT &src, int diagonal=0, const typename base_types::packing<typename RT::dtype>::unpacked_type &val=0) {
apply(dst, src, [val, diagonal]__device__(int row, int col, auto &src_val) {
return col <= row + diagonal ? src_val : val;
});
}
template<ducks::rt::all RT>
__device__ static inline void triu(RT &dst, const RT &src, int diagonal=0, const typename base_types::packing<typename RT::dtype>::unpacked_type &val=0) {
apply(dst, src, [val, diagonal]__device__(int row, int col, auto &src_val) {
return col >= row + diagonal ? src_val : val;
});
}
@@ -1,554 +0,0 @@
/**
* @file
* @brief Reduction operations mapping tiles to vectors.
*/
/**
* @brief Perform a row-wise reduction on a matrix in row-major layout.
*
* This function template performs a parallel reduction across the rows of a matrix using a specified operation.
* It leverages warp shuffle functions for efficient intra-warp communication.
*
* @tparam op The operation to be applied for reduction.
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type with row layout.
* @tparam reset A boolean flag indicating whether to reset the accumulator (ignore src_accum) or not.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when reset is false.
*/
template<typename op, ducks::rv::all V, ducks::rt::row_layout T, bool reset>
__device__ static inline void row_reduce(V &row_accum, const T &src, const V &src_accum) {
// I actually like these static asserts because they give more verbose errors when things go wrong.
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::col_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::height); // compatible size
using dtype = V::dtype;
const int leader = threadIdx.x & 0x1C; // 11100 in binary
#pragma unroll
for(int i = 0; i < src.height; i++) {
dtype accum_top_row = op::template op<dtype>(src.tiles[i][0].data[0], src.tiles[i][0].data[2]);
dtype accum_bottom_row = op::template op<dtype>(src.tiles[i][0].data[1], src.tiles[i][0].data[3]);
#pragma unroll
for(int j = 1; j < src.width; j++) {
#pragma unroll
for(int k = 0; k < src.packed_per_tile; k+=2) {
accum_top_row = op::template op<dtype>(accum_top_row, src.tiles[i][j].data[k+0]);
accum_bottom_row = op::template op<dtype>(accum_bottom_row, src.tiles[i][j].data[k+1]);
}
}
dtype accum_packed;
accum_packed.x = op::template op<typename base_types::packing<dtype>::unpacked_type>(accum_top_row.x, accum_top_row.y);
accum_packed.y = op::template op<typename base_types::packing<dtype>::unpacked_type>(accum_bottom_row.x, accum_bottom_row.y);
// Now we need to do a lil shuffle to make everyone happy.
accum_packed = op::template op<dtype>(accum_packed, packed_shfl_down_sync(MASK_ALL, accum_packed, 2));
accum_packed = op::template op<dtype>(accum_packed, packed_shfl_down_sync(MASK_ALL, accum_packed, 1));
accum_packed = packed_shfl_sync(MASK_ALL, accum_packed, leader);
if(reset) {
row_accum[i][0] = accum_packed;
}
else {
row_accum[i][0] = op::template op<dtype>(src_accum[i][0], accum_packed);
}
}
}
/**
* @brief Perform a row-wise reduction on a matrix in column-major layout.
*
* This function template performs a parallel reduction across the rows of a matrix using a specified operation.
* It leverages warp shuffle functions for efficient intra-warp communication and is optimized for column-major matrices.
*
* @tparam op The operation to be applied for reduction.
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type with column layout.
* @tparam reset A boolean flag indicating whether to reset the accumulator (ignore src_accum) or not.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when reset is false.
*/
template<typename op, ducks::rv::all V, ducks::rt::col_layout T, bool reset>
__device__ static inline void row_reduce(V &row_accum, const T &src, const V &src_accum) {
// I actually like these static asserts because they give more verbose errors when things go wrong.
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::col_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::height); // compatible size
using dtype = V::dtype;
const int leader = threadIdx.x & 0x3; // 00011 in binary
#pragma unroll
for(int i = 0; i < src.height; i++) {
dtype accum_top_rows = op::template op<dtype>(src.tiles[i][0].data[0], src.tiles[i][0].data[1]);
dtype accum_bottom_rows = op::template op<dtype>(src.tiles[i][0].data[2], src.tiles[i][0].data[3]);
#pragma unroll
for(int j = 1; j < src.width; j++) {
#pragma unroll
for(int k = 0; k < src.packed_per_tile/2; k++) {
accum_top_rows = op::template op<dtype>(accum_top_rows, src.tiles[i][j].data[k+0]);
accum_bottom_rows = op::template op<dtype>(accum_bottom_rows, src.tiles[i][j].data[k+2]);
}
}
// Now we need to do a lil shuffle to make everyone happy.
accum_top_rows = op::template op<dtype>(accum_top_rows, packed_shfl_down_sync(MASK_ALL, accum_top_rows, 16));
accum_top_rows = op::template op<dtype>(accum_top_rows, packed_shfl_down_sync(MASK_ALL, accum_top_rows, 8));
accum_top_rows = op::template op<dtype>(accum_top_rows, packed_shfl_down_sync(MASK_ALL, accum_top_rows, 4));
accum_bottom_rows = op::template op<dtype>(accum_bottom_rows, packed_shfl_down_sync(MASK_ALL, accum_bottom_rows, 16));
accum_bottom_rows = op::template op<dtype>(accum_bottom_rows, packed_shfl_down_sync(MASK_ALL, accum_bottom_rows, 8));
accum_bottom_rows = op::template op<dtype>(accum_bottom_rows, packed_shfl_down_sync(MASK_ALL, accum_bottom_rows, 4));
accum_top_rows = packed_shfl_sync(MASK_ALL, accum_top_rows, leader);
accum_bottom_rows = packed_shfl_sync(MASK_ALL, accum_bottom_rows, leader);
if(reset) {
row_accum[i][0] = accum_top_rows;
row_accum[i][1] = accum_bottom_rows;
}
else {
row_accum[i][0] = op::template op<dtype>(src_accum[i][0], accum_top_rows);
row_accum[i][1] = op::template op<dtype>(src_accum[i][1], accum_bottom_rows);
}
}
}
// Col reduction.
/**
* @brief Perform a column-wise reduction on a matrix in row-major layout.
*
* This function template performs a parallel reduction across the columns of a matrix using a specified operation.
* It leverages warp shuffle functions for efficient intra-warp communication and is optimized for row-major matrices.
*
* @tparam op The operation to be applied for reduction.
* @tparam V The vector type for the column accumulator.
* @tparam T The matrix type with row layout.
* @tparam reset A boolean flag indicating whether to reset the accumulator (ignore src_accum) or not.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when reset is false.
*/
template<typename op, ducks::rv::all V, ducks::rt::row_layout T, bool reset>
__device__ static inline void col_reduce(V &col_accum, const T &src, const V &src_accum) {
// I actually like these static asserts because they give more verbose errors when things go wrong.
KITTENS_CHECK_WARP
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::row_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::width); // compatible size
using dtype = V::dtype;
const int leader = threadIdx.x & 0x3; // 00011 in binary
#pragma unroll
for(int j = 0; j < src.width; j++) {
dtype accum_left_cols = op::template op<dtype>(src.tiles[0][j].data[0], src.tiles[0][j].data[1]);
dtype accum_right_cols = op::template op<dtype>(src.tiles[0][j].data[2], src.tiles[0][j].data[3]);
#pragma unroll
for(int i = 1; i < src.height; i++) {
#pragma unroll
for(int k = 0; k < src.packed_per_tile/2; k++) {
accum_left_cols = op::template op<dtype>(accum_left_cols, src.tiles[i][j].data[k+0]);
accum_right_cols = op::template op<dtype>(accum_right_cols, src.tiles[i][j].data[k+2]);
}
}
// Now we need to do a lil shuffle to make everyone happy.
accum_left_cols = op::template op<dtype>(accum_left_cols, packed_shfl_down_sync(MASK_ALL, accum_left_cols, 16));
accum_left_cols = op::template op<dtype>(accum_left_cols, packed_shfl_down_sync(MASK_ALL, accum_left_cols, 8));
accum_left_cols = op::template op<dtype>(accum_left_cols, packed_shfl_down_sync(MASK_ALL, accum_left_cols, 4));
accum_right_cols = op::template op<dtype>(accum_right_cols, packed_shfl_down_sync(MASK_ALL, accum_right_cols, 16));
accum_right_cols = op::template op<dtype>(accum_right_cols, packed_shfl_down_sync(MASK_ALL, accum_right_cols, 8));
accum_right_cols = op::template op<dtype>(accum_right_cols, packed_shfl_down_sync(MASK_ALL, accum_right_cols, 4));
accum_left_cols = packed_shfl_sync(MASK_ALL, accum_left_cols, leader);
accum_right_cols = packed_shfl_sync(MASK_ALL, accum_right_cols, leader);
if(reset) {
col_accum[j][0] = accum_left_cols;
col_accum[j][1] = accum_right_cols;
}
else {
col_accum[j][0] = op::template op<dtype>(src_accum[j][0], accum_left_cols);
col_accum[j][1] = op::template op<dtype>(src_accum[j][1], accum_right_cols);
}
}
}
/**
* @brief Perform a column-wise reduction on a matrix in column-major layout.
*
* This function template performs a parallel reduction across the columns of a matrix using a specified operation.
* It leverages warp shuffle functions for efficient intra-warp communication and is optimized for column-major matrices.
*
* @tparam op The operation to be applied for reduction.
* @tparam V The vector type for the column accumulator.
* @tparam T The matrix type with column layout.
* @tparam reset A boolean flag indicating whether to reset the accumulator (ignore src_accum) or not.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when reset is false.
*/
template<typename op, ducks::rv::all V, ducks::rt::col_layout T, bool reset>
__device__ static inline void col_reduce(V &col_accum, const T &src, const V &src_accum) {
// I actually like these static asserts because they give more verbose errors when things go wrong.
KITTENS_CHECK_WARP
static_assert(std::is_same_v<typename V::layout, typename rt_base<typename T::T, typename T::layout>::row_vec_layout>); // compatible layout
static_assert(std::is_same_v<typename V::dtype, typename T::dtype>); // compatible type
static_assert(V::outer_dim == T::width); // compatible size
using dtype = V::dtype;
const int leader = threadIdx.x & 0x1C; // 11100 in binary
#pragma unroll
for(int j = 0; j < src.width; j++) { // note now width is the outer loop
dtype accum_left_col = op::template op<dtype>(src.tiles[0][j].data[0], src.tiles[0][j].data[2]);
dtype accum_right_col = op::template op<dtype>(src.tiles[0][j].data[1], src.tiles[0][j].data[3]);
#pragma unroll
for(int i = 1; i < src.height; i++) { // and height is the inner loop
#pragma unroll
for(int k = 0; k < src.packed_per_tile; k+=2) {
accum_left_col = op::template op<dtype>(accum_left_col, src.tiles[i][j].data[k+0]);
accum_right_col = op::template op<dtype>(accum_right_col, src.tiles[i][j].data[k+1]);
}
}
dtype accum_packed;
accum_packed.x = op::template op<typename base_types::packing<dtype>::unpacked_type>(accum_left_col.x, accum_left_col.y);
accum_packed.y = op::template op<typename base_types::packing<dtype>::unpacked_type>(accum_right_col.x, accum_right_col.y);
// Now we need to do a lil shuffle to make everyone happy.
accum_packed = op::template op<dtype>(accum_packed, packed_shfl_down_sync(MASK_ALL, accum_packed, 2));
accum_packed = op::template op<dtype>(accum_packed, packed_shfl_down_sync(MASK_ALL, accum_packed, 1));
accum_packed = packed_shfl_sync(MASK_ALL, accum_packed, leader);
if(reset) {
col_accum[j][0] = accum_packed;
}
else {
col_accum[j][0] = op::template op<dtype>(src_accum[j][0], accum_packed);
}
}
}
/* ---------- WRAPPERS FOR PRETTINESS ---------- */
// two-operand row reductions. (Accumulate and REPLACE.)
/**
* @brief Store the maximum of each row of the src register tile in the row_accum column vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void row_max(V &row_accum, const T &src) {
row_reduce<base_ops::max, V, T, true>(row_accum, src, row_accum);
}
/**
* @brief Store the minimum of each row of the src register tile in the row_accum column vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void row_min(V &row_accum, const T &src) {
row_reduce<base_ops::min, V, T, true>(row_accum, src, row_accum);
}
/**
* @brief Store the sum of each row of the src register tile in the row_accum column vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void row_sum(V &row_accum, const T &src) {
row_reduce<base_ops::sum, V, T, true>(row_accum, src, row_accum);
}
/**
* @brief Store the product of each row of the src register tile in the row_accum column vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void row_prod(V &row_accum, const T &src) {
row_reduce<base_ops::mul, V, T, true>(row_accum, src, row_accum);
}
// three-operand row reductions. (Accumulate ONTO.)
/**
* @brief Store the maximum of each row of the src register tile, as well as the src_accum column vector, in the row_accum column vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when accumulating onto an existing value.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void row_max(V &row_accum, const T &src, const V &src_accum) {
row_reduce<base_ops::max, V, T, false>(row_accum, src, src_accum);
}
/**
* @brief Store the minimum of each row of the src register tile, as well as the src_accum column vector, in the row_accum column vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when accumulating onto an existing value.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void row_min(V &row_accum, const T &src, const V &src_accum) {
row_reduce<base_ops::min, V, T, false>(row_accum, src, src_accum);
}
/**
* @brief Store the sum of each row of the src register tile, as well as the src_accum column vector, in the row_accum column vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when accumulating onto an existing value.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void row_sum(V &row_accum, const T &src, const V &src_accum) {
row_reduce<base_ops::sum, V, T, false>(row_accum, src, src_accum);
}
/**
* @brief Store the product of each row of the src register tile, as well as the src_accum column vector, in the row_accum column vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] row_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when accumulating onto an existing value.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void row_prod(V &row_accum, const T &src, const V &src_accum) {
row_reduce<base_ops::mul, V, T, false>(row_accum, src, src_accum);
}
// two-operand col reductions. (Accumulate and REPLACE.)
/**
* @brief Store the maximum of each column of the src register tile in the col_accum row vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void col_max(V &col_accum, const T &src) {
col_reduce<base_ops::max, V, T, true>(col_accum, src, col_accum);
}
/**
* @brief Store the minimum of each column of the src register tile in the col_accum row vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void col_min(V &col_accum, const T &src) {
col_reduce<base_ops::min, V, T, true>(col_accum, src, col_accum);
}
/**
* @brief Store the sum of each column of the src register tile in the col_accum row vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void col_sum(V &col_accum, const T &src) {
col_reduce<base_ops::sum, V, T, true>(col_accum, src, col_accum);
}
/**
* @brief Store the product of each column of the src register tile in the col_accum row vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void col_prod(V &col_accum, const T &src) {
col_reduce<base_ops::mul, V, T, true>(col_accum, src, col_accum);
}
// three-operand col reductions. (Accumulate ONTO.)
/**
* @brief Store the maximum of each column of the src register tile, as well as the src_accum row vector, in the col_accum row vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when accumulating onto an existing value.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void col_max(V &col_accum, const T &src, const V &src_accum) {
col_reduce<base_ops::max, V, T, false>(col_accum, src, src_accum);
}
/**
* @brief Store the minimum of each column of the src register tile, as well as the src_accum row vector, in the col_accum row vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when accumulating onto an existing value.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void col_min(V &col_accum, const T &src, const V &src_accum) {
col_reduce<base_ops::min, V, T, false>(col_accum, src, src_accum);
}
/**
* @brief Store the sum of each column of the src register tile, as well as the src_accum row vector, in the col_accum row vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when accumulating onto an existing value.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void col_sum(V &col_accum, const T &src, const V &src_accum) {
col_reduce<base_ops::sum, V, T, false>(col_accum, src, src_accum);
}
/**
* @brief Store the product of each column of the src register tile, as well as the src_accum row vector, in the col_accum row vector.
*
* @tparam V The vector type for the row accumulator.
* @tparam T The matrix type.
* @param[out] col_accum The accumulator where the result of the reduction is stored.
* @param[in] src The source matrix on which to perform the reduction.
* @param[in] src_accum The initial value of the accumulator, used when accumulating onto an existing value.
*/
template<ducks::rv::all V, ducks::rt::all T>
__device__ static inline void col_prod(V &col_accum, const T &src, const V &src_accum) {
col_reduce<base_ops::mul, V, T, false>(col_accum, src, src_accum);
}
// templated versions of each
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline void max(RV &dst, const T &src, const RV &src_accum) {
if constexpr (ax == axis::COL) row_max(dst, src, src_accum);
else col_max(dst, src, src_accum);
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline auto max(const T &src, const RV &src_accum) {
RV dst;
if constexpr (ax == axis::COL) row_max(dst, src, src_accum);
else col_max(dst, src, src_accum);
return dst;
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline void max(RV &dst, const T &src) {
if constexpr (ax == axis::COL) row_max(dst, src);
else col_max(dst, src);
}
template<int ax, ducks::rt::all T>
__device__ static inline auto max(const T &src) {
using RV = std::conditional_t<ax==axis::COL, typename T::col_vec, typename T::row_vec>;
RV dst;
if constexpr (ax == axis::COL) row_max(dst, src);
else col_max(dst, src);
return dst;
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline void min(RV &dst, const T &src, const RV &src_accum) {
if constexpr (ax == axis::COL) row_min(dst, src, src_accum);
else col_min(dst, src, src_accum);
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline auto min(const T &src, const RV &src_accum) {
RV dst;
if constexpr (ax == axis::COL) row_min(dst, src, src_accum);
else col_min(dst, src, src_accum);
return dst;
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline void min(RV &dst, const T &src) {
if constexpr (ax == axis::COL) row_min(dst, src);
else col_min(dst, src);
}
template<int ax, ducks::rt::all T>
__device__ static inline auto min(const T &src) {
using RV = std::conditional_t<ax==axis::COL, typename T::col_vec, typename T::row_vec>;
RV dst;
if constexpr (ax == axis::COL) row_min(dst, src);
else col_min(dst, src);
return dst;
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline void sum(RV &dst, const T &src, const RV &src_accum) {
if constexpr (ax == axis::COL) row_sum(dst, src, src_accum);
else col_sum(dst, src, src_accum);
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline auto sum(const T &src, const RV &src_accum) {
RV dst;
if constexpr (ax == axis::COL) row_sum(dst, src, src_accum);
else col_sum(dst, src, src_accum);
return dst;
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline void sum(RV &dst, const T &src) {
if constexpr (ax == axis::COL) row_sum(dst, src);
else col_sum(dst, src);
}
template<int ax, ducks::rt::all T>
__device__ static inline auto sum(const T &src) {
using RV = std::conditional_t<ax==axis::COL, typename T::col_vec, typename T::row_vec>;
RV dst;
if constexpr (ax == axis::COL) row_sum(dst, src);
else col_sum(dst, src);
return dst;
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline void prod(RV &dst, const T &src, const RV &src_accum) {
if constexpr (ax == axis::COL) row_prod(dst, src, src_accum);
else col_prod(dst, src, src_accum);
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline auto prod(const T &src, const RV &src_accum) {
RV dst;
if constexpr (ax == axis::COL) row_prod(dst, src, src_accum);
else col_prod(dst, src, src_accum);
return dst;
}
template<int ax, ducks::rv::all RV, ducks::rt::all T>
__device__ static inline void prod(RV &dst, const T &src) {
if constexpr (ax == axis::COL) row_prod(dst, src);
else col_prod(dst, src);
}
template<int ax, ducks::rt::all T>
__device__ static inline auto prod(const T &src) {
using RV = std::conditional_t<ax==axis::COL, typename T::col_vec, typename T::row_vec>;
RV dst;
if constexpr (ax == axis::COL) row_prod(dst, src);
else col_prod(dst, src);
return dst;
}
@@ -1,47 +0,0 @@
/**
* @file
* @brief An aggregate header for warp operations on register tiles.
*/
#include "conversions.cuh"
#include "maps.cuh"
#include "reductions.cuh"
template<ducks::rt::all RT>
__device__ static inline bool hasnan(const RT &src) {
KITTENS_CHECK_WARP
bool nan_detected = false;
#pragma unroll
for(int i = 0; i < RT::height; i++) {
#pragma unroll
for(int j = 0; j < RT::width; j++) {
#pragma unroll
for(int k = 0; k < RT::packed_per_tile; k++) {
if constexpr (std::is_same_v<typename RT::T, float>) {
if(isnan(src.tiles[i][j].data[k].x) || isnan(src.tiles[i][j].data[k].y)) {
nan_detected = true;
}
}
else if constexpr (std::is_same_v<typename RT::T, bf16>) {
if(isnan(__bfloat162float(src.tiles[i][j].data[k].x)) || isnan(__bfloat162float(src.tiles[i][j].data[k].y))) {
nan_detected = true;
}
}
else if constexpr (std::is_same_v<typename RT::T, half>) {
if(isnan(__half2float(src.tiles[i][j].data[k].x)) || isnan(__half2float(src.tiles[i][j].data[k].y))) {
nan_detected = true;
}
}
else {
static_assert(sizeof(typename RT::T) == 999, "Unsupported dtype");
}
}
}
}
// Ballot across the warp to see if any lane detected a nan
return (__ballot_sync(0xffffffff, nan_detected) != 0);
}
#include "complex/complex_conversions.cuh"
#include "complex/complex_maps.cuh"
@@ -1,153 +0,0 @@
/**
* @file
* @brief Conversions on vectors stored in registers.
*/
struct vec_conversion_detail {
// i am not smart enough to figure out these indices without these helpers :/
// again, blame nvidia for these stupid, stupid layouts
__device__ static inline int row_from_indices_dim2(int laneid, int inner_dim, int x_or_y) {
return 8*inner_dim + (laneid%4)*2 + x_or_y;
}
__device__ static inline int row_from_indices_dim1(int laneid, int x_or_y) {
return 8*x_or_y + (laneid/4);
}
__device__ static inline int canonical_src_lane_dim2(int row) {
return (row/2)%4 + 4*(row%2); // draw even rows from 0...3 and odds from 4...7
}
__device__ static inline int canonical_src_lane_dim1(int row) {
return (row*4)%32;
}
};
/**
* @brief Copies data from one register vector to another.
*
* @tparam RV1 The type of the destination register vector.
* @tparam RV2 The type of the source register vector.
* @param dst[out] The destination register vector.
* @param src[in] The source register vector to copy from.
*/
template<ducks::rv::all RV1, ducks::rv::all RV2>
__device__ static inline void copy(RV1 &dst, const RV2 &src) {
KITTENS_CHECK_WARP
static_assert(RV1::length == RV2::length, "Register vectors must be the same length.");
using D1 = RV1::dtype;
using D2 = RV2::dtype;
if constexpr (std::is_same_v<typename RV1::layout, typename RV2::layout>) { // just a simple copy / typecast
#pragma unroll
for(int i = 0; i < RV1::outer_dim; i++) {
#pragma unroll
for(int j = 0; j < RV1::inner_dim; j++) {
dst[i][j] = base_types::convertor<D1, D2>::convert(src[i][j]);
}
}
}
else { // Inner dimensions are not the same, this is really a layout conversion.
int laneid = ::kittens::laneid();
if constexpr (std::is_same_v<typename RV1::layout, ortho_l> && std::is_same_v<typename RV2::layout, align_l>) { // align -> ortho layout
#pragma unroll
for(int i = 0; i < RV1::outer_dim; i++) {
dst[i][0].x = packed_shfl_sync(
kittens::MASK_ALL,
laneid < 4 ? src[i][0].x : src[i][0].y, // mirrors canonical_src_lane_dim2
vec_conversion_detail::canonical_src_lane_dim2(vec_conversion_detail::row_from_indices_dim1(laneid, 0))
);
dst[i][0].y = packed_shfl_sync(
kittens::MASK_ALL,
laneid < 4 ? src[i][1].x : src[i][1].y, // mirrors canonical_src_lane_dim2
vec_conversion_detail::canonical_src_lane_dim2(vec_conversion_detail::row_from_indices_dim1(laneid, 1))
);
}
}
else if constexpr (std::is_same_v<typename RV1::layout, align_l> && std::is_same_v<typename RV2::layout, ortho_l>) { // ortho -> align layout
#pragma unroll
for(int i = 0; i < RV1::outer_dim; i++) {
dst[i][0].x = packed_shfl_sync(
kittens::MASK_ALL,
src[i][0].x, // first 8 rows
vec_conversion_detail::canonical_src_lane_dim1(vec_conversion_detail::row_from_indices_dim2(laneid, 0, 0))
);
dst[i][0].y = packed_shfl_sync(
kittens::MASK_ALL,
src[i][0].x, // first 8 rows
vec_conversion_detail::canonical_src_lane_dim1(vec_conversion_detail::row_from_indices_dim2(laneid, 0, 1))
);
dst[i][1].x = packed_shfl_sync(
kittens::MASK_ALL,
src[i][0].y, // last 8 rows
vec_conversion_detail::canonical_src_lane_dim1(vec_conversion_detail::row_from_indices_dim2(laneid, 1, 0))
);
dst[i][1].y = packed_shfl_sync(
kittens::MASK_ALL,
src[i][0].y, // last 8 rows
vec_conversion_detail::canonical_src_lane_dim1(vec_conversion_detail::row_from_indices_dim2(laneid, 1, 1))
);
}
}
else if constexpr (std::is_same_v<typename RV1::layout, ortho_l> && std::is_same_v<typename RV2::layout, naive_l>) { // naive -> ortho layout
#pragma unroll
for(int i = 0; i < RV1::outer_dim; i++) {
dst[i][0].x = packed_shfl_sync(
kittens::MASK_ALL, src[i/2][0],
16*(i%2) + 0 + (laneid/4)
);
dst[i][0].y = packed_shfl_sync(
kittens::MASK_ALL, src[i/2][0],
16*(i%2) + 8 + (laneid/4)
);
}
}
else if constexpr (std::is_same_v<typename RV1::layout, naive_l> && std::is_same_v<typename RV2::layout, ortho_l>) { // ortho -> naive layout
int lane_replication = laneid%4; // 0...3
#pragma unroll
for(int i = 0; i < RV1::outer_dim; i++) {
D1 tmp = 0;
if(RV1::length%32==0 || i < RV1::outer_dim-1 || lane_replication<2) {
tmp = lane_replication%2 ? src[2*i + (lane_replication>=2)][0].y : src[2*i + (lane_replication>=2)][0].x;
}
dst[i][0] = packed_shfl_sync(
kittens::MASK_ALL, tmp,
(laneid%8)*4 + (laneid/8)
);
}
}
else if constexpr (std::is_same_v<typename RV1::layout, align_l> && std::is_same_v<typename RV2::layout, naive_l>) { // naive -> align layout
#pragma unroll
for(int i = 0; i < RV1::outer_dim; i++) {
dst[i][0].x = packed_shfl_sync(
kittens::MASK_ALL, src[i/2][0],
16*(i%2) + 0 + 2*(laneid%4) + 0
);
dst[i][0].y = packed_shfl_sync(
kittens::MASK_ALL, src[i/2][0],
16*(i%2) + 0 + 2*(laneid%4) + 1
);
dst[i][1].x = packed_shfl_sync(
kittens::MASK_ALL, src[i/2][0],
16*(i%2) + 8 + 2*(laneid%4) + 0
);
dst[i][1].y = packed_shfl_sync(
kittens::MASK_ALL, src[i/2][0],
16*(i%2) + 8 + 2*(laneid%4) + 1
);
}
}
else if constexpr (std::is_same_v<typename RV1::layout, naive_l> && std::is_same_v<typename RV2::layout, align_l>) { // align -> naive layout
int lane_replication = laneid/8; // 0...3
#pragma unroll
for(int i = 0; i < RV1::outer_dim; i++) {
D1 tmp = 0;
if(RV1::length%32==0 || i < RV1::outer_dim-1 || laneid<16) {
tmp = (laneid%8)<4 ? src[2*i + (lane_replication>=2)][lane_replication%2].x : src[2*i + (lane_replication>=2)][lane_replication%2].y;
}
dst[i][0] = packed_shfl_sync(
kittens::MASK_ALL, tmp,
4*(laneid%2) + (laneid%8)/2 + (laneid&0b11000)
);
}
}
}
}
@@ -1,374 +0,0 @@
/**
* @file
* @brief Maps on vectors stored in registers.
*/
/* ---------- Vector Maps ---------- */
/**
* @brief Perform a unary operation on a vector.
*
* @tparam op The unary operation to perform.
* @tparam T The type of the vector.
* @param dst[out] The destination vector where the result is stored.
* @param src[in] The source vector to perform the operation on.
*/
template<typename op, ducks::rv::all T>
__device__ static inline void unary_op(T &dst, const T &src) {
#pragma unroll
for(int i = 0; i < dst.outer_dim; i++) {
#pragma unroll
for(int j = 0; j < dst.inner_dim; j++) {
dst[i][j] = op::template op<typename T::dtype>(src[i][j]);
}
}
}
/**
* @brief Perform a binary operation on two vectors.
*
* @tparam op The binary operation to perform.
* @tparam T The type of the vectors.
* @param dst[out] The destination vector where the result is stored.
* @param lhs[in] The left-hand side vector for the operation.
* @param rhs[in] The right-hand side vector for the operation.
*/
template<typename op, ducks::rv::all T>
__device__ static inline void bin_op(T &dst, const T &lhs, const T &rhs) {
#pragma unroll
for(int i = 0; i < dst.outer_dim; i++) {
#pragma unroll
for(int j = 0; j < dst.inner_dim; j++) {
dst[i][j] = op::template op<typename T::dtype>(lhs[i][j], rhs[i][j]);
}
}
}
/**
* @brief Perform a binary operation on a vector and a scalar.
*
* @tparam op The binary operation to perform.
* @tparam T The type of the vector.
* @param dst[out] The destination vector where the result is stored.
* @param src[in] The source vector for the operation.
* @param param[in] The scalar parameter for the operation.
*/
template<typename op, ducks::rv::all T>
__device__ static inline void bin_op(T &dst, const T &src, const typename T::dtype &param) {
#pragma unroll
for(int i = 0; i < dst.outer_dim; i++) {
#pragma unroll
for(int j = 0; j < dst.inner_dim; j++) {
dst[i][j] = op::template op<typename T::dtype>(src[i][j], param);
}
}
}
/**
* @brief Perform a binary operation on a vector and an unpacked scalar.
*
* @tparam op The binary operation to perform.
* @tparam T The type of the vector.
* @param dst[out] The destination vector where the result is stored.
* @param src[in] The source vector for the operation.
* @param param[in] The unpacked scalar parameter for the operation.
*/
template<typename op, ducks::rv::tile_layout T>
__device__ static inline void bin_op(T &dst, const T &src, const typename base_types::packing<typename T::dtype>::unpacked_type &param) {
bin_op<op, T>(dst, src, base_types::packing<typename T::dtype>::pack(param));
}
template<ducks::rv::all RV, typename Lambda>
__device__ static inline void apply(RV &dst, const RV &src, Lambda &&lambda) {
int group_offset = 0;
if constexpr(GROUP_WARPS > 1) {
group_offset = warpid()*RV::length;
}
static_assert(sizeof(RV::T) != 1, "Cannot apply lambda to 8-bit types");
if constexpr (ducks::rv::ortho_layout<RV>) {
#pragma unroll
for(int i = 0; i < dst.outer_dim; i++) {
int base_idx = group_offset + i*16 + ::kittens::laneid()/4;
dst[i][0].x = lambda(base_idx+0, src[i][0].x);
dst[i][0].y = lambda(base_idx+8, src[i][0].y);
}
}
else if constexpr (ducks::rv::align_layout<RV>) {
#pragma unroll
for(int i = 0; i < dst.outer_dim; i++) {
int base_idx = group_offset + i*16 + 2*(::kittens::laneid()%4);
dst[i][0].x = lambda(base_idx+0, src[i][0].x);
dst[i][0].y = lambda(base_idx+1, src[i][0].y);
dst[i][1].x = lambda(base_idx+8, src[i][1].x);
dst[i][1].y = lambda(base_idx+9, src[i][1].y);
}
}
else {
#pragma unroll
for(int i = 0; i < dst.outer_dim; i++) {
int base_idx = group_offset + i*32 + ::kittens::laneid();
if (i < dst.outer_dim-1 || dst.length%32 == 0 || ::kittens::laneid()<16) {
dst[i][0] = lambda(base_idx, src[i][0]);
}
}
}
}
template<ducks::rv::all RV, typename Lambda>
__device__ static inline RV apply(const RV &src, Lambda &&lambda) {
RV dst;
apply<RV, Lambda>(dst, src, std::forward<Lambda>(lambda));
return dst;
}
/* ---------- WRAPPERS FOR PRETTINESS ---------- */
// ---- const ops ----
/**
* @brief Sets all elements of a register vector to zero.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector to be set to zero.
*/
template<ducks::rv::all T>
__device__ static inline void zero(T &dst) {
unary_op<base_ops::zero, T>(dst, dst);
}
/**
* @brief Sets all elements of a register vector to one.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector to be set to one.
*/
template<ducks::rv::all T>
__device__ static inline void one(T &dst) {
unary_op<base_ops::one, T>(dst, dst);
}
/**
* @brief Sets all elements of a register vector to positive infinity.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector to be set to positive infinity.
*/
template<ducks::rv::all T>
__device__ static inline void pos_infty(T &dst) {
unary_op<base_ops::pos_infty, T>(dst, dst);
}
/**
* @brief Sets all elements of a register vector to negative infinity.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector to be set to negative infinity.
*/
template<ducks::rv::all T>
__device__ static inline void neg_infty(T &dst) {
unary_op<base_ops::neg_infty, T>(dst, dst);
}
// ---- unary ops ----
/**
* @brief Copies the elements from one register vector to another.
*
* @tparam T Register vector type.
* @tparam U Type of the source vector.
* @param dst[out] Destination vector where the elements will be copied to.
* @param src[in] Source vector to copy the elements from.
*/
template<ducks::rv::all T, typename U>
__device__ static inline void copy(T &dst, const U &src) {
bin_op<base_ops::copy2, T>(dst, dst, src); // the second arg is ignored here.
}
/**
* @brief Applies the exponential function element-wise to a register vector.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector where the exponential values will be stored.
* @param src[in] Source vector to apply the exponential function to.
*/
template<ducks::rv::all T>
__device__ static inline void exp(T &dst, const T &src) {
unary_op<base_ops::exp, T>(dst, src);
}
template<ducks::rv::all T>
__device__ static inline T exp(const T &src) {
T dst;
exp(dst, src);
return dst;
}
/**
* @brief Applies the exponential function element-wise to a register vector, in base 2.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector where the exponential values will be stored.
* @param src[in] Source vector to apply the exponential function to.
*/
template<ducks::rv::all T>
__device__ static inline void exp2(T &dst, const T &src) {
unary_op<base_ops::exp2, T>(dst, src);
}
template<ducks::rv::all T>
__device__ static inline T exp2(const T &src) {
T dst;
exp2(dst, src);
return dst;
}
/**
* @brief Applies the natural logarithm function element-wise to a register vector.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector where the exponential values will be stored.
* @param src[in] Source vector to apply the exponential function to.
*/
template<ducks::rv::all T>
__device__ static inline void log(T &dst, const T &src) {
unary_op<base_ops::log, T>(dst, src);
}
template<ducks::rv::all T>
__device__ static inline T log(const T &src) {
T dst;
log(dst, src);
return dst;
}
/**
* @brief Applies the logarithm base 2 function element-wise to a register vector.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector where the exponential values will be stored.
* @param src[in] Source vector to apply the logarithm base 2 function to.
*/
template<ducks::rv::all T>
__device__ static inline void log2(T &dst, const T &src) {
unary_op<base_ops::log2, T>(dst, src);
}
template<ducks::rv::all T>
__device__ static inline T log2(const T &src) {
T dst;
log2(dst, src);
return dst;
}
/**
* @brief Applies the absolute value function element-wise to a register vector.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector where the absolute values will be stored.
* @param src[in] Source vector to apply the absolute value function to.
*/
template<ducks::rv::all T>
__device__ static inline void abs(T &dst, const T &src) {
unary_op<base_ops::abs, T>(dst, src);
}
template<ducks::rv::all T>
__device__ static inline T abs(const T &src) {
T dst;
abs(dst, src);
return dst;
}
/**
* @brief Applies the rectified linear unit (ReLU) function element-wise to a register vector.
*
* @tparam T Register vector type.
* @param dst[out] Destination vector where the ReLU values will be stored.
* @param src[in] Source vector to apply the ReLU function to.
*/
template<ducks::rv::all T>
__device__ static inline void relu(T &dst, const T &src) {
unary_op<base_ops::relu, T>(dst, src);
}
template<ducks::rv::all T>
__device__ static inline T relu(const T &src) {
T dst;
relu(dst, src);
return dst;
}
// ---- binary ops ----
/**
* @brief Computes the element-wise maximum of two register vectors.
*
* @tparam T Register vector type.
* @tparam U Type of the second vector.
* @param dst[out] Destination vector where the maximum values will be stored.
* @param lhs[in] First vector for the maximum operation.
* @param rhs[in] Second vector for the maximum operation.
*/
template<ducks::rv::all T, typename U>
__device__ static inline void max(T &dst, const T &lhs, const U &rhs) {
bin_op<base_ops::max, T>(dst, lhs, rhs);
}
template<ducks::rv::all T, typename U>
__device__ static inline T max(const T &lhs, const U &rhs) {
T dst;
max(dst, lhs, rhs);
return dst;
}
/**
* @brief Computes the element-wise minimum of two register vectors.
*
* @tparam T Register vector type.
* @tparam U Type of the second vector.
* @param dst[out] Destination vector where the minimum values will be stored.
* @param lhs[in] First vector for the minimum operation.
* @param rhs[in] Second vector for the minimum operation.
*/
template<ducks::rv::all T, typename U>
__device__ static inline void min(T &dst, const T &lhs, const U &rhs) {
bin_op<base_ops::min, T>(dst, lhs, rhs);
}
template<ducks::rv::all T, typename U>
__device__ static inline T min(const T &lhs, const U &rhs) {
T dst;
min(dst, lhs, rhs);
return dst;
}
/**
* @brief Computes the element-wise sum of two register vectors.
*
* @tparam T Register vector type.
* @tparam U Type of the second vector.
* @param dst[out] Destination vector where the sum values will be stored.
* @param lhs[in] First vector for the sum operation.
* @param rhs[in] Second vector for the sum operation.
*/
template<ducks::rv::all T, typename U>
__device__ static inline void add(T &dst, const T &lhs, const U &rhs) {
bin_op<base_ops::sum, T>(dst, lhs, rhs);
}
/**
* @brief Computes the element-wise difference of two register vectors.
*
* @tparam T Register vector type.
* @tparam U Type of the second vector.
* @param dst[out] Destination vector where the difference values will be stored.
* @param lhs[in] First vector for the difference operation.
* @param rhs[in] Second vector for the difference operation.
*/
template<ducks::rv::all T, typename U>
__device__ static inline void sub(T &dst, const T &lhs, const U &rhs) {
bin_op<base_ops::sub, T>(dst, lhs, rhs);
}
/**
* @brief Computes the element-wise product of two register vectors.
*
* @tparam T Register vector type.
* @tparam U Type of the second vector.
* @param dst[out] Destination vector where the product values will be stored.
* @param lhs[in] First vector for the product operation.
* @param rhs[in] Second vector for the product operation.
*/
template<ducks::rv::all T, typename U>
__device__ static inline void mul(T &dst, const T &lhs, const U &rhs) {
bin_op<base_ops::mul, T>(dst, lhs, rhs);
}
/**
* @brief Computes the element-wise division of two register vectors.
*
* @tparam T Register vector type.
* @tparam U Type of the second vector.
* @param dst[out] Destination vector where the division values will be stored.
* @param lhs[in] First vector for the division operation.
* @param rhs[in] Second vector for the division operation.
*/
template<ducks::rv::all T, typename U>
__device__ static inline void div(T &dst, const T &lhs, const U &rhs) {
bin_op<base_ops::div, T>(dst, lhs, rhs);
}

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