From 53e1b2f4e879aa486a31b7ee05f636afbaaaa649 Mon Sep 17 00:00:00 2001 From: Vincent Ouyang Date: Fri, 9 Oct 2026 16:50:15 -0700 Subject: [PATCH] Upgrade MI355X runtime and evaluation tools to ROCm 10 Pin the new scoring image, adapt SDK loading and task dependencies, and update all six evaluation-tool runtimes with matching documentation and regression coverage. Keep qualification pending for GPU ASan and the existing fused-MoE numerical failures. --- README.md | 2 +- agents/geak/README.md | 6 +- docker/eval-tools/gpu-asan/Dockerfile | 52 +- docker/eval-tools/gpu-asan/packages.sha256 | 8 +- docker/eval-tools/hip-fpsan/Dockerfile | 4 +- docker/eval-tools/images.lock.yaml | 51 +- .../eval-tools/rocjitsu-sanitizers/Dockerfile | 16 +- docker/eval-tools/rocjitsu/Dockerfile | 11 +- docker/eval-tools/triton-fpsan/Dockerfile | 19 +- .../eval-tools/triton-fpsan/requirements.lock | 4 - docs/how-to/use-evaluation-tools.md | 73 ++- docs/index.rst | 1 + docs/install/install.md | 44 +- docs/reference/api-reference.md | 10 +- docs/reference/compatibility-matrix.md | 26 +- docs/reference/mi355x-runtime.md | 62 ++ docs/reference/release-notes.md | 5 +- docs/sphinx/_toc.yml.in | 2 + .../evaluation_tools_advisory_mi355x.yaml | 4 +- src/eval_tools/README.md | 12 +- src/eval_tools/adapters/flydsl_aot.py | 8 +- src/eval_tools/adapters/triton_aot.py | 40 +- src/eval_tools/config.py | 2 + src/eval_tools/contracts.py | 2 +- src/eval_tools/plugins/gpu_asan.py | 23 +- src/eval_tools/runtime_client.py | 1 + src/eval_tools/worker.py | 65 +- src/scripts/docker_benchmark.sh | 45 +- src/scripts/rocm_sdk_runtime.py | 52 ++ .../README.md | 4 + .../flydsl_compat/LICENSE | 17 + .../flydsl_compat/SOURCE.md | 13 + .../flydsl_compat/__init__.py | 13 + .../flydsl_compat/buffer_ops.py | 603 ++++++++++++++++++ .../flydsl_compat/meta.py | 65 ++ .../flydsl_compat/vector.py | 121 ++++ .../kernel.py | 24 +- .../kernels/kernels_common.py | 2 +- .../pa_decode_fp8_kernel/README.md | 11 +- .../pa_decode_fp8_kernel/config.yaml | 2 +- .../flydsl_compat/LICENSE | 17 + .../flydsl_compat/SOURCE.md | 13 + .../flydsl_compat/__init__.py | 13 + .../flydsl_compat/buffer_ops.py | 603 ++++++++++++++++++ .../flydsl_compat/meta.py | 65 ++ .../flydsl_compat/vector.py | 121 ++++ .../pa_decode_fp8_kernel/kernel.py | 3 +- .../kernels/pa_decode_swa.py | 3 +- .../batched_gemm_a8w8_kernel/README.md | 11 + .../scripts/task_actions.py | 10 +- .../batched_gemm_a8w8_kernel/task_runtime.py | 14 +- .../gemm_a8w8_bpreshuffle_kernel/README.md | 10 +- .../flydsl_compat/LICENSE | 17 + .../flydsl_compat/SOURCE.md | 13 + .../flydsl_compat/__init__.py | 13 + .../flydsl_compat/buffer_ops.py | 603 ++++++++++++++++++ .../flydsl_compat/meta.py | 65 ++ .../flydsl_compat/vector.py | 121 ++++ .../gemm_a8w8_bpreshuffle_kernel/kernel.py | 11 +- .../scripts/task_actions.py | 10 +- tasks/torch2flydsl/hgemm_kernel/README.md | 32 +- .../hgemm_kernel/flydsl_compat/LICENSE | 17 + .../hgemm_kernel/flydsl_compat/SOURCE.md | 13 + .../hgemm_kernel/flydsl_compat/__init__.py | 13 + .../hgemm_kernel/flydsl_compat/buffer_ops.py | 603 ++++++++++++++++++ .../hgemm_kernel/flydsl_compat/meta.py | 65 ++ .../hgemm_kernel/flydsl_compat/vector.py | 121 ++++ tasks/torch2flydsl/hgemm_kernel/kernel.py | 3 +- .../hgemm_kernel/scripts/task_actions.py | 10 +- .../torch2flydsl/hgemm_kernel/task_runtime.py | 51 +- .../jagged_dense_bmm_kernel/README.md | 7 +- .../flydsl_compat/LICENSE | 17 + .../flydsl_compat/SOURCE.md | 13 + .../flydsl_compat/__init__.py | 13 + .../flydsl_compat/buffer_ops.py | 603 ++++++++++++++++++ .../flydsl_compat/meta.py | 65 ++ .../flydsl_compat/vector.py | 121 ++++ .../jagged_dense_bmm_kernel/kernel.py | 23 +- .../torch2flydsl/moe_sorting_kernel/README.md | 5 + .../moe_sorting_kernel/flydsl_compat/LICENSE | 17 + .../flydsl_compat/SOURCE.md | 13 + .../flydsl_compat/__init__.py | 13 + .../flydsl_compat/buffer_ops.py | 603 ++++++++++++++++++ .../moe_sorting_kernel/flydsl_compat/meta.py | 65 ++ .../flydsl_compat/vector.py | 121 ++++ .../torch2flydsl/moe_sorting_kernel/kernel.py | 3 +- .../qk_norm_rope_quant_kernel/README.md | 23 +- .../qk_norm_rope_quant_kernel/config.yaml | 4 +- .../flydsl_compat/LICENSE | 17 + .../flydsl_compat/SOURCE.md | 13 + .../flydsl_compat/__init__.py | 13 + .../flydsl_compat/buffer_ops.py | 603 ++++++++++++++++++ .../flydsl_compat/meta.py | 65 ++ .../flydsl_compat/vector.py | 121 ++++ .../qk_norm_rope_quant_kernel/kernel.py | 5 +- tasks/triton2flydsl/aiter/mla/README.md | 12 +- tasks/triton2flydsl/aiter/mla/mla.py | 4 +- tests/eval_tools/test_manager.py | 15 +- tests/eval_tools/test_plugins.py | 34 + tests/eval_tools/test_replay_capsule.py | 6 +- tests/eval_tools/test_runtime_client.py | 2 + tests/eval_tools/test_triton_aot.py | 56 ++ tests/eval_tools/test_worker.py | 76 +++ tests/test_docker_benchmark.sh | 62 +- tests/test_flydsl_gfx950_vendor_compat.py | 12 +- tests/test_flydsl_task_migration_v2.py | 103 ++- tests/test_mla_runtime_compat.py | 40 ++ tests/test_rocm_sdk_runtime.py | 46 ++ tests/test_task_materialization_v2.py | 21 +- 109 files changed, 6973 insertions(+), 271 deletions(-) delete mode 100644 docker/eval-tools/triton-fpsan/requirements.lock create mode 100644 docs/reference/mi355x-runtime.md create mode 100755 src/scripts/rocm_sdk_runtime.py create mode 100644 tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/LICENSE create mode 100644 tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/SOURCE.md create mode 100644 tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/__init__.py create mode 100644 tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/buffer_ops.py create mode 100644 tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/meta.py create mode 100644 tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/vector.py create mode 100644 tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/LICENSE create mode 100644 tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/SOURCE.md create mode 100644 tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/__init__.py create mode 100644 tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/buffer_ops.py create mode 100644 tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/meta.py create mode 100644 tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/vector.py create mode 100644 tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/LICENSE create mode 100644 tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/SOURCE.md create mode 100644 tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/__init__.py create mode 100644 tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/buffer_ops.py create mode 100644 tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/meta.py create mode 100644 tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/vector.py create mode 100644 tasks/torch2flydsl/hgemm_kernel/flydsl_compat/LICENSE create mode 100644 tasks/torch2flydsl/hgemm_kernel/flydsl_compat/SOURCE.md create mode 100644 tasks/torch2flydsl/hgemm_kernel/flydsl_compat/__init__.py create mode 100644 tasks/torch2flydsl/hgemm_kernel/flydsl_compat/buffer_ops.py create mode 100644 tasks/torch2flydsl/hgemm_kernel/flydsl_compat/meta.py create mode 100644 tasks/torch2flydsl/hgemm_kernel/flydsl_compat/vector.py create mode 100644 tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/LICENSE create mode 100644 tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/SOURCE.md create mode 100644 tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/__init__.py create mode 100644 tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/buffer_ops.py create mode 100644 tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/meta.py create mode 100644 tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/vector.py create mode 100644 tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/LICENSE create mode 100644 tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/SOURCE.md create mode 100644 tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/__init__.py create mode 100644 tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/buffer_ops.py create mode 100644 tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/meta.py create mode 100644 tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/vector.py create mode 100644 tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/LICENSE create mode 100644 tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/SOURCE.md create mode 100644 tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/__init__.py create mode 100644 tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/buffer_ops.py create mode 100644 tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/meta.py create mode 100644 tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/vector.py create mode 100644 tests/eval_tools/test_triton_aot.py create mode 100644 tests/test_mla_runtime_compat.py create mode 100644 tests/test_rocm_sdk_runtime.py diff --git a/README.md b/README.md index cbe60d55f..897d150b8 100755 --- a/README.md +++ b/README.md @@ -167,7 +167,7 @@ The prompt system also recognizes `cuda2hip`; the current bundled task tree does - Node.js 22+ and npm when using the alternative npm installation of Claude Code (or another npm-installed agent CLI); DeepSeek Harness uses a dedicated Node.js 24 prefix as described in [its setup guide](agents/deepseek_harness/README.md) -- For MI300/MI355X, use the GPU-specific SGLang image: `gfx942` uses `lmsysorg/sglang:v0.5.12-rocm720-mi30x`; `gfx950` uses `lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705` +- For MI300/MI355X, use the GPU-specific SGLang image: `gfx942` uses `lmsysorg/sglang:v0.5.12-rocm720-mi30x`; `gfx950` uses the digest-pinned SGLang 0.5.20 / ROCm 10 runtime. See the [runtime compatibility matrix](docs/reference/compatibility-matrix.md) for image selection and task-specific requirements. - For RDNA4 `gfx1201`, the runner automatically builds the default [pinned RDNA4 runtime](docker/rdna4/README.md) on first use if it is missing. - A supported agent CLI installed and logged in on the host, or the dependencies required by a specialized agent diff --git a/agents/geak/README.md b/agents/geak/README.md index e720de8c0..8ab4d53c6 100644 --- a/agents/geak/README.md +++ b/agents/geak/README.md @@ -63,7 +63,11 @@ Docker provisions Claude for `geak` and its v2 aliases, mounts only the selected GEAK checkout read-only, and forwards `GEAK_HOME`. The complete checkout is needed for revision checks and private engine/knowledge copies. Preflight installs the pinned SDK into the GEAK-only dependency directory when necessary -and verifies the clean upstream pin. These checks do not certify a live +and verifies the clean upstream pin. SDK dependencies are cached under +`.aka-pyuserbase/geak-sdk/` so changing the image's Python version +installs compatible wheels. The first run after an upgrade may reinstall the +SDK; existing caches remain available for their original interpreter. +These checks do not certify a live Workflow invocation or a GPU task. ## Upstream compatibility and evaluation diff --git a/docker/eval-tools/gpu-asan/Dockerfile b/docker/eval-tools/gpu-asan/Dockerfile index a3c2bf17d..399d0a4e8 100644 --- a/docker/eval-tools/gpu-asan/Dockerfile +++ b/docker/eval-tools/gpu-asan/Dockerfile @@ -1,30 +1,48 @@ -ARG BASE_IMAGE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 +ARG BASE_IMAGE=lmsysorg/sglang:v0.5.20-rocm10-mi35x@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69 FROM ${BASE_IMAGE} LABEL org.opencontainers.image.title="AgentKernelArena ROCm GPU ASan tool runtime" \ - org.opencontainers.image.base.digest="sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78" \ - org.opencontainers.image.version="rocm-7.2.0-asan" + org.opencontainers.image.base.digest="sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69" \ + org.opencontainers.image.version="rocm-10.0.0-asan" COPY docker/eval-tools/gpu-asan/packages.sha256 /tmp/packages.sha256 -RUN mkdir -p /tmp/gpu-asan-debs \ +# Unpack checksum-locked ROCm 10 ASan packages into their own versioned prefix. +# Installing their metapackage would also mutate the base image's SDK selection. +RUN mkdir -p /tmp/gpu-asan-debs /tmp/gpu-asan-root \ && cd /tmp/gpu-asan-debs \ - && apt-get update \ - && apt-get download \ - rocm-core-asan=7.2.0.70200-43~22.04 \ - comgr-asan=3.0.0.70200-43~22.04 \ - hsa-rocr-asan=1.18.0.70200-43~22.04 \ - hip-runtime-amd-asan=7.2.26015.70200-43~22.04 \ + && while read -r checksum package; do \ + curl --fail --location --retry 3 -o "$package" \ + "https://stable.repo.amd.com/rocm/core/packages-asan/ubuntu2404/pool/main/$package"; \ + done < /tmp/packages.sha256 \ && sha256sum --check /tmp/packages.sha256 \ - && DEBIAN_FRONTEND=noninteractive apt-get install -y --no-install-recommends ./*.deb \ - && rm -rf /tmp/gpu-asan-debs /tmp/packages.sha256 /var/lib/apt/lists/* + && for package in *.deb; do dpkg-deb --extract "$package" /tmp/gpu-asan-root; done \ + && mv /tmp/gpu-asan-root/opt/rocm/core-asan-10.0 /opt/rocm-asan-10.0 \ + && test -L /opt/rocm \ + && test -x /opt/rocm/bin/hipcc \ + && rm -rf /tmp/gpu-asan-debs /tmp/gpu-asan-root /tmp/packages.sha256 + +# PyTorch preloads these SDK libraries by absolute path. Point those paths at +# the instrumented copies too; LD_LIBRARY_PATH alone would load two LLVMs. +RUN /opt/venv/bin/python - <<'PY' +from pathlib import Path +import rocm_sdk + +runtime = Path('/opt/rocm-asan-10.0/lib') +for name in ('amd_comgr', 'amdhip64'): + target = runtime / f'lib{name}.so' + assert target.is_file(), target + for path in rocm_sdk.find_libraries(name): + path.unlink() + path.symlink_to(target) +PY COPY src/eval_tools /opt/aka-eval-tools/src/eval_tools ENV AKA_EVAL_TOOL_FRAMEWORK_ROOT=/opt/aka-eval-tools \ PYTHONPATH=/opt/aka-eval-tools \ - AKA_GPU_ASAN_RUNTIME_DIR=/opt/rocm-7.2.0/lib/asan \ - AKA_GPU_ASAN_HIP_RUNTIME=/opt/rocm-7.2.0/lib/asan/libamdhip64.so \ - AKA_GPU_ASAN_HOST_PRELOAD=/opt/rocm-7.2.0/lib/llvm/lib/clang/22/lib/linux/libclang_rt.asan-x86_64.so \ - AKA_GPU_ASAN_HOST_LIB_DIR=/opt/rocm-7.2.0/lib/llvm/lib/clang/22/lib/linux \ - AKA_GPU_ASAN_NORMAL_ROCM_LIB_DIR=/opt/rocm-7.2.0/lib + AKA_GPU_ASAN_RUNTIME_DIR=/opt/rocm-asan-10.0/lib \ + AKA_GPU_ASAN_HIP_RUNTIME=/opt/rocm-asan-10.0/lib/libamdhip64.so \ + AKA_GPU_ASAN_HOST_PRELOAD=/opt/rocm-asan-10.0/lib/llvm/lib/clang/23/lib/linux/libclang_rt.asan-x86_64.so \ + AKA_GPU_ASAN_HOST_LIB_DIR=/opt/rocm-asan-10.0/lib/llvm/lib/clang/23/lib/linux \ + AKA_GPU_ASAN_NORMAL_ROCM_LIB_DIR=/opt/rocm/lib diff --git a/docker/eval-tools/gpu-asan/packages.sha256 b/docker/eval-tools/gpu-asan/packages.sha256 index 6d8078c4f..ee0d265b7 100644 --- a/docker/eval-tools/gpu-asan/packages.sha256 +++ b/docker/eval-tools/gpu-asan/packages.sha256 @@ -1,4 +1,4 @@ -3bd5b98b3ae2cb8fbfd10c248682feb86d0f5136914e7a226b701340f9b90f83 rocm-core-asan_7.2.0.70200-43~22.04_amd64.deb -31118ea2dc79fe9d8c69ad7ff2176a2f1c822128ece604d4438b6c13a1aa3179 comgr-asan_3.0.0.70200-43~22.04_amd64.deb -c5d6e48846f6163b5c8d7168949e1c24dc41f69bd871705621b0972c2f1a03cc hsa-rocr-asan_1.18.0.70200-43~22.04_amd64.deb -fa56c2192f28adb022dc323d2836dfbe4f12211287575ea27978dc96b10a5300 hip-runtime-amd-asan_7.2.26015.70200-43~22.04_amd64.deb +030fc18229683369c77e7c9e124dd1b80114e7708a5ba8e86b985b031c2cda96 amdrocm-base-asan10.0_10.0.0-4_amd64.deb +3f45a0f26b4c9c01741fea8d42dec83f16dae4f754253d09fba32f1c32aec724 amdrocm-llvm-asan10.0_10.0.0-4_amd64.deb +c23ff5640c087457980d0d8e9edda49b3eb1c2114dd21fe437ffcf5e942fd15b amdrocm-runtime-asan10.0_10.0.0-4_amd64.deb +ff3c7cecf40d8133692116455266a7c0c02a2bdf5b695fdae746e47e50ff7147 amdrocm-sysdeps-asan10.0_10.0.0-4_amd64.deb diff --git a/docker/eval-tools/hip-fpsan/Dockerfile b/docker/eval-tools/hip-fpsan/Dockerfile index d7b4a4e96..8794a3afb 100644 --- a/docker/eval-tools/hip-fpsan/Dockerfile +++ b/docker/eval-tools/hip-fpsan/Dockerfile @@ -1,10 +1,10 @@ -ARG BASE_IMAGE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 +ARG BASE_IMAGE=lmsysorg/sglang:v0.5.20-rocm10-mi35x@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69 FROM ${BASE_IMAGE} ARG HIP_FPSAN_COMMIT=0ac9be8a1539a473ba21dfa686564c3be33c890e LABEL org.opencontainers.image.title="AgentKernelArena HIP-FpSan tool runtime" \ - org.opencontainers.image.base.digest="sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78" \ + org.opencontainers.image.base.digest="sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69" \ org.opencontainers.image.revision="0ac9be8a1539a473ba21dfa686564c3be33c890e" COPY src/eval_tools/probes/hip_fpsan_probe.hip /tmp/hip_fpsan_probe.hip diff --git a/docker/eval-tools/images.lock.yaml b/docker/eval-tools/images.lock.yaml index 90a243ab8..1c80d587e 100644 --- a/docker/eval-tools/images.lock.yaml +++ b/docker/eval-tools/images.lock.yaml @@ -1,36 +1,41 @@ schema_version: 1 -# The scoring runtime remains unchanged. Every tool image is an immutable child -# of this verified MI355X/gfx950 image. +# Every tool runtime is a child of the same pinned MI355X/gfx950 scoring base. +# Native rocJITsu build dependencies have a separate builder lock below. base: gfx950: - reference: lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705 - digest: sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 - rocm: 7.2.0 + reference: lmsysorg/sglang:v0.5.20-rocm10-mi35x + digest: sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69 + rocm: 10.0.0 + +# The native C++ engines keep their GCC 13/Ubuntu 22.04 build toolchain. +# Only installed engines and their C++ runtime libraries enter the ROCm 10 +# sidecars; HIP candidate/control compilation uses the new runtime SDK. +builders: + rocjitsu: + reference: lmsysorg/sglang-rocm@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 tools: triton_fpsan: triton: - version: 3.7.0+amd.rocm7.2.0.gitd0d77a509 - sha256: 3a8acacdfb4723c8bb71844c427f4d3ce047658bf97fbdc23065c06205231687 - triton_kernels: - version: 1.0.0+amd.rocm7.2.0.gitd0d77a509 - sha256: df8b42ebf098767c0d31a3916913bd7557a623903aefa26602e7c6818c34b84a + version: 3.8.0+git4cff872c.rocm10.0.0 + source: pinned_base_image gpu_asan: + repository: https://stable.repo.amd.com/rocm/core/packages-asan/ubuntu2404 packages: - rocm-core-asan: - version: 7.2.0.70200-43~22.04 - sha256: 3bd5b98b3ae2cb8fbfd10c248682feb86d0f5136914e7a226b701340f9b90f83 - comgr-asan: - version: 3.0.0.70200-43~22.04 - sha256: 31118ea2dc79fe9d8c69ad7ff2176a2f1c822128ece604d4438b6c13a1aa3179 - hsa-rocr-asan: - version: 1.18.0.70200-43~22.04 - sha256: c5d6e48846f6163b5c8d7168949e1c24dc41f69bd871705621b0972c2f1a03cc - hip-runtime-amd-asan: - version: 7.2.26015.70200-43~22.04 - sha256: fa56c2192f28adb022dc323d2836dfbe4f12211287575ea27978dc96b10a5300 + amdrocm-base-asan10.0: + version: 10.0.0-4 + sha256: 030fc18229683369c77e7c9e124dd1b80114e7708a5ba8e86b985b031c2cda96 + amdrocm-llvm-asan10.0: + version: 10.0.0-4 + sha256: 3f45a0f26b4c9c01741fea8d42dec83f16dae4f754253d09fba32f1c32aec724 + amdrocm-runtime-asan10.0: + version: 10.0.0-4 + sha256: c23ff5640c087457980d0d8e9edda49b3eb1c2114dd21fe437ffcf5e942fd15b + amdrocm-sysdeps-asan10.0: + version: 10.0.0-4 + sha256: ff3c7cecf40d8133692116455266a7c0c02a2bdf5b695fdae746e47e50ff7147 rocjitsu: repository: https://github.com/ROCm/rocm-systems.git @@ -70,5 +75,5 @@ tools: image_probe: src/eval_tools/probes/hip_fpsan_probe.hip verification: - gfx950: verified + gfx950: qualification_pending gfx942: unverified diff --git a/docker/eval-tools/rocjitsu-sanitizers/Dockerfile b/docker/eval-tools/rocjitsu-sanitizers/Dockerfile index 9178b765c..58033fec2 100644 --- a/docker/eval-tools/rocjitsu-sanitizers/Dockerfile +++ b/docker/eval-tools/rocjitsu-sanitizers/Dockerfile @@ -1,5 +1,8 @@ -ARG BASE_IMAGE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 -FROM ${BASE_IMAGE} AS builder +ARG BASE_IMAGE=lmsysorg/sglang:v0.5.20-rocm10-mi35x@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69 +# Keep the native engine build toolchain pinned independently of its runtime. +ARG BUILDER_IMAGE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 +FROM ${BUILDER_IMAGE} AS builder +ARG BUILD_JOBS=8 ARG ROCJITSU_SANITIZERS_COMMIT=ed35c0b54547c98bab359c8732529d9f5e8fd1ae ARG GOOGLETEST_COMMIT=b514bdc898e2951020cbdca1304b75f5950d1f59 @@ -45,7 +48,7 @@ RUN curl -sSfL -o /tmp/zstd.tar.gz \ -DZSTD_BUILD_STATIC=ON \ -DZSTD_BUILD_PROGRAMS=OFF \ -DZSTD_BUILD_TESTS=OFF \ - && cmake --build /build/zstd --parallel \ + && cmake --build /build/zstd --parallel "${BUILD_JOBS}" \ && cmake --install /build/zstd \ && rm -f /tmp/zstd.tar.gz @@ -58,7 +61,7 @@ RUN cmake -S /src/rocm-systems/emulation/rocjitsu -B /build/rocjitsu -G Ninja \ -DFETCHCONTENT_SOURCE_DIR_GOOGLETEST=/src/googletest \ -DFETCHCONTENT_SOURCE_DIR_FLATBUFFERS=/src/flatbuffers \ -DBUILD_TESTING=OFF \ - && cmake --build /build/rocjitsu --parallel \ + && cmake --build /build/rocjitsu --parallel "${BUILD_JOBS}" \ --target rj_waitcheck rocjitsu_waitcheck rocjitsu_dbi_hooks \ && mkdir -p /opt/rocjitsu/include /opt/rocjitsu/lib \ && cp -a /src/rocm-systems/emulation/rocjitsu/lib/rocjitsu/include/rocjitsu \ @@ -81,7 +84,7 @@ RUN g++-13 -std=c++17 -O2 -Wall -Wextra -Werror \ FROM ${BASE_IMAGE} AS runtime-common -LABEL org.opencontainers.image.base.digest="sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78" \ +LABEL org.opencontainers.image.base.digest="sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69" \ org.opencontainers.image.revision="ed35c0b54547c98bab359c8732529d9f5e8fd1ae" COPY --from=builder /opt/rocjitsu /opt/rocjitsu @@ -89,7 +92,7 @@ COPY src/eval_tools /opt/aka-eval-tools/src/eval_tools ENV AKA_EVAL_TOOL_FRAMEWORK_ROOT=/opt/aka-eval-tools \ PYTHONPATH=/opt/aka-eval-tools \ - LD_LIBRARY_PATH=/opt/rocjitsu/lib:/opt/rocjitsu/runtime \ + LD_LIBRARY_PATH=/opt/rocjitsu/lib:/opt/rocjitsu/runtime:${LD_LIBRARY_PATH} \ PATH=/opt/rocjitsu/bin:${PATH} \ AKA_ROCJITSU_SANITIZERS_COMMIT=ed35c0b54547c98bab359c8732529d9f5e8fd1ae @@ -112,6 +115,7 @@ LABEL org.opencontainers.image.title="AgentKernelArena rocJITsu ConSan runtime" ENV AKA_CONSAN_HOOK=/opt/rocjitsu/lib/librocjitsu_dbi_hooks.so RUN test -f /opt/rocjitsu/lib/librocjitsu_dbi_hooks.so \ + && /opt/venv/bin/python -c "import ctypes; ctypes.CDLL('libamdhip64.so.7')" \ && /opt/venv/bin/python -I \ /opt/aka-eval-tools/src/eval_tools/adapters/consan_entrypoint.py \ --help >/dev/null diff --git a/docker/eval-tools/rocjitsu/Dockerfile b/docker/eval-tools/rocjitsu/Dockerfile index 7a0c2c0ae..ded2a9915 100644 --- a/docker/eval-tools/rocjitsu/Dockerfile +++ b/docker/eval-tools/rocjitsu/Dockerfile @@ -1,5 +1,8 @@ -ARG BASE_IMAGE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 -FROM ${BASE_IMAGE} AS builder +ARG BASE_IMAGE=lmsysorg/sglang:v0.5.20-rocm10-mi35x@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69 +# Keep the native engine build toolchain pinned independently of its runtime. +ARG BUILDER_IMAGE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 +FROM ${BUILDER_IMAGE} AS builder +ARG BUILD_JOBS=8 ARG ROCJITSU_COMMIT=0bf561a0d8a4a6b88954f2c46bd3a50871cda140 ARG GOOGLETEST_COMMIT=b514bdc898e2951020cbdca1304b75f5950d1f59 @@ -40,7 +43,7 @@ RUN cmake -S /src/rocm-systems/emulation/rocjitsu -B /build/rocjitsu -G Ninja \ -DCMAKE_HIP_ARCHITECTURES=gfx950 \ -DBUILD_TESTING=ON \ -DRJ_INSTALL_TESTS=OFF \ - && cmake --build /build/rocjitsu --parallel \ + && cmake --build /build/rocjitsu --parallel "${BUILD_JOBS}" \ && cmake --install /build/rocjitsu \ && mkdir -p /opt/rocjitsu/runtime \ && cp -a /usr/lib/x86_64-linux-gnu/libstdc++.so.6* /opt/rocjitsu/runtime/ \ @@ -49,7 +52,7 @@ RUN cmake -S /src/rocm-systems/emulation/rocjitsu -B /build/rocjitsu -G Ninja \ FROM ${BASE_IMAGE} LABEL org.opencontainers.image.title="AgentKernelArena rocJITsu tool runtime" \ - org.opencontainers.image.base.digest="sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78" \ + org.opencontainers.image.base.digest="sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69" \ org.opencontainers.image.revision="0bf561a0d8a4a6b88954f2c46bd3a50871cda140" COPY --from=builder /opt/rocjitsu /opt/rocjitsu diff --git a/docker/eval-tools/triton-fpsan/Dockerfile b/docker/eval-tools/triton-fpsan/Dockerfile index fa7819cc8..fe43b241c 100644 --- a/docker/eval-tools/triton-fpsan/Dockerfile +++ b/docker/eval-tools/triton-fpsan/Dockerfile @@ -1,21 +1,16 @@ -ARG BASE_IMAGE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 +ARG BASE_IMAGE=lmsysorg/sglang:v0.5.20-rocm10-mi35x@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69 FROM ${BASE_IMAGE} LABEL org.opencontainers.image.title="AgentKernelArena Triton FpSan tool runtime" \ - org.opencontainers.image.base.digest="sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78" \ - org.opencontainers.image.version="triton-3.7.0+amd.rocm7.2.0.gitd0d77a509" + org.opencontainers.image.base.digest="sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69" \ + org.opencontainers.image.version="triton-3.8.0+git4cff872c.rocm10.0.0" -COPY docker/eval-tools/triton-fpsan/requirements.lock /tmp/eval-tool-requirements.lock - -RUN python -m pip uninstall -y \ - triton pytorch-triton pytorch-triton-rocm triton-rocm amd-triton triton-kernels \ - && python -m pip install --no-cache-dir --no-deps --require-hashes \ - -r /tmp/eval-tool-requirements.lock \ - && python -c "import importlib.metadata as m, triton, triton_kernels; assert m.version('triton') == '3.7.0+amd.rocm7.2.0.gitd0d77a509'; assert m.version('triton-kernels') == '1.0.0+amd.rocm7.2.0.gitd0d77a509'" \ - && rm -f /tmp/eval-tool-requirements.lock +# The digest-pinned ROCm 10 base already supplies FpSan-enabled Triton. +# Do not replace it with a ROCm 7.2 wheel or infer support from import success. +RUN python -c "import importlib.metadata as m; assert m.version('triton') == '3.8.0+git4cff872c.rocm10.0.0'" COPY src/eval_tools /opt/aka-eval-tools/src/eval_tools ENV AKA_EVAL_TOOL_FRAMEWORK_ROOT=/opt/aka-eval-tools \ PYTHONPATH=/opt/aka-eval-tools \ - AKA_TRITON_FPSAN_VERSION=3.7.0+amd.rocm7.2.0.gitd0d77a509 + AKA_TRITON_FPSAN_VERSION=3.8.0+git4cff872c.rocm10.0.0 diff --git a/docker/eval-tools/triton-fpsan/requirements.lock b/docker/eval-tools/triton-fpsan/requirements.lock deleted file mode 100644 index d8fc54608..000000000 --- a/docker/eval-tools/triton-fpsan/requirements.lock +++ /dev/null @@ -1,4 +0,0 @@ -triton @ https://pypi.amd.com/triton/release_/rocm-7.2.0/packages/triton/triton-3.7.0+amd.rocm7.2.0.gitd0d77a509-cp310-cp310-linux_x86_64.whl \ - --hash=sha256:3a8acacdfb4723c8bb71844c427f4d3ce047658bf97fbdc23065c06205231687 -triton-kernels @ https://pypi.amd.com/triton/release_/rocm-7.2.0/packages/triton-kernels/triton_kernels-1.0.0+amd.rocm7.2.0.gitd0d77a509-py3-none-any.whl \ - --hash=sha256:df8b42ebf098767c0d31a3916913bd7557a623903aefa26602e7c6818c34b84a diff --git a/docs/how-to/use-evaluation-tools.md b/docs/how-to/use-evaluation-tools.md index a90890706..f68d58e7a 100644 --- a/docs/how-to/use-evaluation-tools.md +++ b/docs/how-to/use-evaluation-tools.md @@ -24,14 +24,16 @@ initial tool set is: This feature is experimental and opt-in. Capability and evidence checks fail closed; whether an incomplete result blocks performance is controlled by the `advisory` or `required` policy. It is not a general sanitizer suite: -sanitizer suite: every result is qualified by the kernel language, generated +every result is qualified by the kernel language, generated artifact, adapter, tool image, GPU architecture, and evidence that the intended kernel was actually instrumented or dispatched. > **Current validation boundary:** sidecar build locks, integrated startup -> controls, and end-to-end fixtures exist only for MI355X (`gfx950`). All six -> startup controls passed in the current hardware qualification. Candidate -> readiness still depends on language, artifact, adapter, and attestation. +> controls, and end-to-end fixtures exist only for MI355X (`gfx950`). The ROCm 10 +> migration has passed five runtime startup controls; GPU ASan qualification is +> blocked by failing safe probes on the tested hosts. See the +> [runtime compatibility guide](../reference/mi355x-runtime.md). +> Candidate readiness still depends on language, artifact, adapter, and attestation. > `gfx942` is unverified, and the Docker runner currently rejects > evaluation-tool sidecars on that architecture. Do not interpret normal > MI300/MI325 task support as sanitizer support. @@ -77,11 +79,11 @@ that it imported this image-owned tree rather than the repository mounted at checkout used by a running worker; candidate-specific commands and inputs remain separate, explicitly mounted data. -The verified `gfx950` scoring image remains: +The ROCm 10 `gfx950` tool profile uses this scoring image: ```text -lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705 -lmsysorg/sglang-rocm@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 +lmsysorg/sglang:v0.5.20-rocm10-mi35x +lmsysorg/sglang@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69 ``` When any evaluation tool is enabled, the runner resolves Docker's immutable @@ -94,10 +96,13 @@ looks compatible. The selected reference and verified image ID are recorded unde `plan.source_evidence.metadata.scoring_runtime`, serialized with the report, and covered by the plan fingerprint. -The current design deliberately does **not** upgrade that image, FlyDSL 0.2.2, -or AITER `0.1.17.dev110+g9127c94a1`. The Triton FpSan sidecar replaces Triton -only inside its own container; the other tool dependencies are likewise local -to their sidecars. A sidecar is not a replacement scoring image and must not be +All six sidecars use the same digest-pinned ROCm 10 base as scoring. Triton +FpSan uses the base image's compiler and checks its exact package version; it +does not install the older Python 3.10 / ROCm 7.2 wheels. GPU ASan installs its +checksum-locked runtime into a separate versioned directory. Native rocJITsu +engines retain their pinned GCC build stage, then run and compile candidate +fixtures in the ROCm 10 runtime stage. Tool-specific dependencies stay inside +their sidecars. A sidecar is not a replacement scoring image and must not be used to establish a new performance baseline. FlyDSL does not promise that every generated artifact or task remains compatible @@ -111,8 +116,8 @@ The pinned sidecar dependencies are recorded in | Sidecar | Isolated dependency change | | --- | --- | -| `triton_fpsan` | AMD Triton `3.7.0+amd.rocm7.2.0.gitd0d77a509` and matching `triton-kernels` wheels. | -| `gpu_asan` | ROCm 7.2 ASan runtime packages, including `hip-runtime-amd-asan`. | +| `triton_fpsan` | Bundled AMD Triton `3.8.0+git4cff872c.rocm10.0.0`; equivalent and known-wrong comparisons each require compiler instrumentation metadata. | +| `gpu_asan` | Checksum-locked ROCm 10 ASan packages, separate from the normal scoring SDK. | | `rocjitsu` | rocJITsu from pinned `rocm-systems` commit `0bf561a0...`, built with GCC 13 for `gfx950`. | | `rocjitsu_waitcheck` | rocJITsu Waitcheck C API and CLI from pinned `rocm-systems` commit `ed35c0b...`; zstd source is separately checksum-locked. | | `rocjitsu_consan` | rocJITsu ConSan HSA hook from the same pinned `ed35c0b...` source, forced to strict record/replay mode. | @@ -171,7 +176,7 @@ it does not mean the ordinary correctness command is automatically reused. | --- | --- | --- | --- | --- | | Editable Triton Python/JIT | Ready with comparison adapter and instrumentation attestation | Ready with dedicated command, fresh JIT cache, XNACK, and build attestation | Trusted `triton_aot` capsule replay is implemented on `gfx950`; whole-Python JIT remains unsupported, and capsule capture/binding to the correctness run is not automatic, so use it only as advisory evidence | Not applicable | | HIP source controlled by the task | Not applicable | Ready only after recompiling the candidate with `-fsanitize=address -shared-libsan --offload-arch=gfx950:xnack+`, then attesting that artifact | Ready with a dedicated native launcher | Source port and comparison adapter required; both reference and candidate paths must explicitly use `fpsan::Value` | -| FlyDSL 0.2.2 Python/JIT | Unsupported; FlyDSL does not use the Triton FpSan pipeline | Unsupported; the current ROCDL pipeline does not insert AMD GPU ASan instrumentation | Trusted `flydsl_aot` capsule replay is implemented on `gfx950` and detects the seeded LDS race; automatic capsule capture/binding to the correctness run is not ready, so use it only as advisory evidence | Not applicable | +| FlyDSL Python/JIT | Unsupported; FlyDSL does not use the Triton FpSan pipeline | Unsupported; the current ROCDL pipeline does not insert AMD GPU ASan instrumentation | Trusted `flydsl_aot` capsule replay is implemented on `gfx950` and detects the seeded LDS race; automatic capsule capture/binding to the correctness run is not ready, so use it only as advisory evidence | Not applicable | | Editable Triton source inside AITER | Engine may be eligible for the explicitly selected source only, with a dedicated comparison adapter; this does not sanitize AITER library kernels | Unsupported by the current default AITER runtime path | Unsupported by the current Python/AITER runtime | Not applicable | | AITER or another precompiled HSACO/library kernel | Cannot retrofit instrumentation | Unsupported unless the exact kernel source is rebuilt and attested; preloading the runtime is insufficient | Unsupported by the current evaluator runtime | Cannot retrofit value semantics | | rocBLAS or RCCL internal kernel | Do not enable; library internals are outside the selected submission | The stock library is not instrumented and is not covered | Not a supported general library-runtime path | Do not enable | @@ -198,7 +203,7 @@ is qualified. The current `gfx950` startup qualification is stricter: | --- | --- | | Triton FpSan | Passing on hardware; eligible task paths can proceed to candidate attestation. | | HIP-FpSan | Passing on hardware; explicitly ported task paths can proceed to candidate attestation. | -| GPU ASan | Passing on hardware for both HIP and Triton safe/OOB lanes; an applicable candidate still needs its own instrumentation/build attestation. | +| GPU ASan | Not qualified on the new runtime: safe probes fail on the tested hosts. The runtime stays unavailable when its required controls fail; historical ROCm 7.2 results do not qualify this image. | | rocJITsu | Passing on hardware with barrier-safe and deliberately racy LDS fixtures; an applicable candidate still needs a native HIP launcher or validated AOT replay capsule. | | rocJITsu Waitcheck | Passing on hardware: a correct `s_waitcnt lgkmcnt(0)` fixture is clean and a missing-wait fixture produces one exact hazard. Candidate use still requires exact SHA-256, kernel name, and entry attestation. | | rocJITsu ConSan | Passing on hardware in strict record/replay: a single-wave LDS fixture is clean and a two-wave conflicting fixture produces complete FNV-attributed diagnostics. Candidate use still requires an exact code object, focused loader, and separate oracle. | @@ -244,12 +249,12 @@ src/scripts/docker_benchmark.sh build-eval-tool-images The default local tags are: ```text -agent-kernel-arena/eval-tool-triton-fpsan:gfx950 -agent-kernel-arena/eval-tool-gpu-asan:gfx950 -agent-kernel-arena/eval-tool-rocjitsu:gfx950 -agent-kernel-arena/eval-tool-rocjitsu-waitcheck:gfx950 -agent-kernel-arena/eval-tool-rocjitsu-consan:gfx950 -agent-kernel-arena/eval-tool-hip-fpsan:gfx950 +agent-kernel-arena/eval-tool-triton-fpsan:gfx950-rocm10 +agent-kernel-arena/eval-tool-gpu-asan:gfx950-rocm10 +agent-kernel-arena/eval-tool-rocjitsu:gfx950-rocm10 +agent-kernel-arena/eval-tool-rocjitsu-waitcheck:gfx950-rocm10 +agent-kernel-arena/eval-tool-rocjitsu-consan:gfx950-rocm10 +agent-kernel-arena/eval-tool-hip-fpsan:gfx950-rocm10 ``` Check that the workers start and report their pinned assets: @@ -279,7 +284,7 @@ worker: | Tool | Startup positive control | | --- | --- | -| `triton_fpsan` | Compile instrumented reference/candidate kernels and require a known numerical mismatch to produce different digests plus FpSan compiler metadata. | +| `triton_fpsan` | Compile both equivalent and known-wrong reference/candidate pairs in separate caches; require matching/distinct digests respectively and FpSan compiler metadata for every compiled kernel. | | `gpu_asan` | Compile and run safe/OOB HIP fixtures and safe/OOB Triton fixtures; the task profile selects the relevant lane. | | `rocjitsu` | Require a barrier-protected fixture to remain clean and a deliberately racy LDS fixture to report a race. | | `rocjitsu_waitcheck` | Compile unbundled `gfx950` code objects and run the production entrypoint, inventory, C API, and parser on the correct-wait and missing-wait fixtures; retain a direct CLI hazard check as an independent engine control. | @@ -292,13 +297,14 @@ JSON summaries before promotion. A normal evaluation with `positive_control: required` repeats the fail-closed check during the typed runtime probe. -As of the current `gfx950` qualification run, all six integrated startup -controls pass on hardware. This qualifies the installed tool runtimes only. It -does not promote a candidate path without the language-specific adapter and +On the new `gfx950` runtime, five integrated startup controls pass on hardware; +GPU ASan remains unqualified. Startup controls qualify only an installed runtime. +They do not promote a candidate path without the language-specific adapter and attestation in the strict support matrix. -The same final image set also passed evaluator-manager-to-sidecar candidate -fixtures on the physical MI355X host: +The previous ROCm 7.2 image set passed these evaluator-manager-to-sidecar +candidate fixtures on a physical MI355X host. These historical results do not +qualify the new ROCm 10 images: | Tool and language | Safe fixture | Seeded bug fixture | | --- | --- | --- | @@ -321,8 +327,8 @@ global milestone. | Phase | Work | Exit criterion | | --- | --- | --- | -| 0. Freeze baselines | Keep the pinned scoring image, FlyDSL 0.2.2, and AITER version unchanged; build each tool from its lock into a sidecar. | Existing compilation, correctness, held-out, and performance baselines remain unchanged with tools disabled. Sidecar image IDs and the verified scoring-image ID/reference are captured in plans. | -| 1. Qualify installations | Run automatic safe/known-bug startup controls on `gfx950`; repeat the now-passing six-tool qualification on clean hosts. | Both positive and negative lanes pass repeatedly. `eval-tools-smoke` evidence is archived and independently reviewed. | +| 0. Freeze baselines | Freeze the selected scoring image and its bundled dependencies for each comparison; build each tool from its lock into a sidecar. | Existing compilation, correctness, held-out, and performance baselines remain unchanged with tools disabled. Sidecar image IDs and the verified scoring-image ID/reference are captured in plans. | +| 1. Qualify installations | Run automatic safe/known-bug startup controls on `gfx950`; require all six tools to pass on compatible hosts. | Both positive and negative lanes pass repeatedly. `eval-tools-smoke` evidence is archived and independently reviewed. | | 2. Build trusted pilot adapters | Start with one editable Triton task for Triton FpSan, one Triton and one HIP task for GPU ASan, one native HIP task for rocJITsu, one final-HSACO task for Waitcheck, one focused native loader for ConSan, and one explicitly ported HIP-FpSan task. Put harnesses under protected `scripts/` paths and declare all inputs. | Each pilot distinguishes a safe fixture from a seeded bug, identifies the selected candidate, and produces bounded structured artifacts. No precompiled AITER/library kernel is claimed as broadly covered. | | 3. Finish AOT capture and binding | The trusted `triton_aot`/`flydsl_aot` replay path now validates one-dispatch capsules and generates the launcher. Add evaluator-owned extraction immediately after correctness and bind the capsule to that exact candidate/case. | Safe and racy fixtures pass end to end, malformed capsules fail closed, and a task cannot substitute a different valid capsule for the correctness dispatch. | | 4. Harden provenance and phase isolation | The runner now uses per-tool writable socket directories, a read-only socket parent in scoring, a narrow per-worker artifact mount, fresh per-invocation artifact directories, a complete serialized plan, and capsule digests in the fingerprint. Next run tools only after the agent exits, freeze the candidate, use evaluator-only/authenticated RPC and evaluator-owned artifacts, strengthen artifact/dispatch binding, and wire resume to plan freshness. | An adversarial task cannot call a worker, overwrite evidence, reach another task's artifacts, spoof a clean result, or reuse a stale report. This phase is required before sanitizer output becomes a reward signal. | @@ -508,7 +514,7 @@ evaluation_tools: oracle_command: [scripts/load_hsaco, build/optimized.hsaco, --check] ``` -With ROCm 7.2, `hipcc --genco` produces a clang bundle by default; use +`hipcc --genco` can produce a clang bundle rather than a raw code object; use `--no-gpu-bundle-output` or explicitly extract the final device ELF before supplying `code_object`. @@ -696,7 +702,7 @@ tool_evaluation: metadata: scoring_runtime: image_id: "sha256:..." - reference: "lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705" + reference: "lmsysorg/sglang:v0.5.20-rocm10-mi35x" policy: advisory overall_status: incomplete resolved_task_profile: {} @@ -987,8 +993,9 @@ the following: runner. - Every useful task still needs a reviewed adapter command. Tool installation alone usually produces `adapter_required`. -- All six startup positive controls pass on the current `gfx950` host. This - qualifies tool installation, not candidate coverage. +- Five startup controls pass for the ROCm 10 profile. GPU ASan remains + unqualified on the tested hosts. Passing controls qualify tool installation, + not candidate coverage. - Runtime-internal asset paths are injected from verified sidecar health and cannot be supplied by task configuration. - Build-attestation artifact paths must be relative to the attestation file; diff --git a/docs/index.rst b/docs/index.rst index 4b3e09bb3..7e5282a44 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -41,6 +41,7 @@ repository. * :doc:`Configuration and API reference ` * :doc:`Performance measurement methodology ` + * :doc:`MI355X runtime and compatibility ` To contribute to the documentation, see the `AgentKernelArena GitHub repository `_. diff --git a/docs/install/install.md b/docs/install/install.md index 635bde0d5..ab1fd126e 100644 --- a/docs/install/install.md +++ b/docs/install/install.md @@ -23,7 +23,8 @@ The following prerequisites are required before running AgentKernelArena. without `sudo`. - **Runtime image:** `gfx942` uses `lmsysorg/sglang:v0.5.12-rocm720-mi30x`; `gfx950` uses - `lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705`. The runner selects from + the digest-pinned SGLang 0.5.20 / ROCm 10 image declared in + [`docker_benchmark.sh`](../../src/scripts/docker_benchmark.sh). The runner selects from `target_gpu_model` for experiment runs and from the visible host GPU for shell and smoke commands. For `gfx1201`, the runner automatically builds the default @@ -57,8 +58,45 @@ before a container is launched. The [DeepSeek draft validator config](../../example_configs/task_validator_deepseek_drafts_mi355x.yaml) pins its MI355X run to the tested ROCm 10 image. The runner supplies writable -AITER/FlyDSL caches and AITER configuration storage for that image. The global -architecture defaults and evaluation-tool image qualification are unchanged. +AITER/FlyDSL caches and AITER configuration storage for that image, which is also +the MI355X default. Other GPU architecture defaults are unchanged. + +Start a separate run after changing images. Ordinary runs record the selected +image reference in materialization and baseline state, and resume rejects a +different runtime. Older workspaces without that image identity also require a +fresh run; their runtime cannot be verified retrospectively. Pin custom images +by digest when comparing experiments. +An image upgrade changes the compiler and runtime libraries as well as ROCm; +measure baseline and candidate in the same runtime and keep older results +associated with their original image. + +To roll back, select `GFX950_V0514_IMMUTABLE_IMAGE` from the +[runner image definitions](../../src/scripts/docker_benchmark.sh) with +`docker_image` or `AKA_DOCKER_IMAGE`, and start a fresh run. Evaluation tools +also need the matching repository revision and sidecar set; the ROCm 10 tool +profile deliberately rejects the old scoring image. + +The MI355X runner selects the SDK core libraries for workloads and routes +`rocprofv3` to that same installation used by PyTorch. +The image's original console entrypoint chooses a separate developer tree; +loading both trees in a Python workload can abort with duplicate LLVM option +registration. The wrapper preserves profiler arguments and GPU selection. +`rocprof-compute` is an optional package and is not bundled in this image; see +the [runtime guide](../reference/mi355x-runtime.md#sdk-and-profiling) for +SDK library selection and profiler availability boundaries. + +Some image-backed tasks require a dedicated runtime, including tasks that +materialize vLLM source directories. Select their documented image explicitly +through the run config's `docker_image` field; the MI355X default does not supply +every task's external dependencies. Optional evaluation tools have their own +[scoring-image compatibility requirements](../how-to/use-evaluation-tools.md). + +To reproduce a run on the previous MI355X runtime, explicitly select its digest: + +```bash +AKA_DOCKER_IMAGE_GFX950=lmsysorg/sglang-rocm@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78 \ + make docker-run CONFIG=example_configs/quickstart_claude_mi355x.yaml +``` ```bash git clone https://github.com/AMD-AGI/AgentKernelArena.git diff --git a/docs/reference/api-reference.md b/docs/reference/api-reference.md index 8c7fb7cb3..8f8c90e3c 100644 --- a/docs/reference/api-reference.md +++ b/docs/reference/api-reference.md @@ -59,17 +59,17 @@ The built-in IDs are `triton_fpsan`, `gpu_asan`, `rocjitsu`, `rocjitsu_waitcheck`, `rocjitsu_consan`, and `hip_fpsan`. Sidecar build locks, integrated positive controls, and end-to-end fixtures -currently exist only for `gfx950`; all six startup controls pass in the current -MI355X qualification. Each applicable candidate still needs a task-specific +currently exist only for `gfx950`. The ROCm 10 migration has five passing +startup controls; GPU ASan remains unqualified, as recorded in the +[runtime compatibility guide](mi355x-runtime.md). Each applicable candidate still needs a task-specific adapter and attestation, and enabling an image alone does not imply that a kernel was analyzed. See [Check kernels with evaluation tools](../how-to/use-evaluation-tools.md) for the support matrix and operational requirements. When tools are enabled, the selected scoring image must resolve to the same -immutable local Docker image ID as the pinned -`lmsysorg/sglang-rocm@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78` -manifest. The runner rejects a different build, launches by the verified ID, +immutable local Docker image ID as the scoring manifest pinned in +[`docker/eval-tools/images.lock.yaml`](../../docker/eval-tools/images.lock.yaml). The runner rejects a different build, launches by the verified ID, and records both the selected reference and verified ID in plan source evidence. Worker reports live at repository-root diff --git a/docs/reference/compatibility-matrix.md b/docs/reference/compatibility-matrix.md index be888a017..4c81f0898 100644 --- a/docs/reference/compatibility-matrix.md +++ b/docs/reference/compatibility-matrix.md @@ -17,7 +17,7 @@ The following hardware configurations are supported and tested. | --- | --- | --- | | MI300X | 7.2 (Bundled in the selected SGLang image.) | `target_gpu_model: MI300X` | | MI325X | 7.2 (Bundled in the selected SGLang image.) | `target_gpu_model: MI325X` | -| MI355X | 7.2 (Bundled in the selected SGLang image.) | `target_gpu_model: MI355X` | +| MI355X | 10.0 (Bundled in the selected SGLang image.) | `target_gpu_model: MI355X` | | RDNA4 (`gfx1201`, 16 GB tested) | Pinned in the [RDNA4 recipe](../../docker/rdna4/Dockerfile) | `target_gpu_model: RDNA4`; HIP/Triton task runtime checks, with limits in the [recipe guide](../../docker/rdna4/README.md). | ## Software requirements @@ -28,28 +28,28 @@ The following software versions are required or verified. | --- | --- | --- | | Linux | Ubuntu 22.04, Ubuntu 24.04 | | | hipcc | Matches ROCm image | Required for HIP tasks. | -| Profiler tools | Match runtime image | Smoke reports `rocprof-compute` and `rocprofv3` availability. Core graph/event timing needs neither; profiling runs can require specific binaries with `AKA_REQUIRED_PROFILERS`. Availability does not establish candidate analysis. See the [qualification record](runtime-upgrade-qualification.md#profiler-capability-policy). | +| Profiler tools | Match runtime image | Smoke reports `rocprof-compute` and `rocprofv3` availability. Core graph/event timing needs neither; profiling runs can require specific binaries with `AKA_REQUIRED_PROFILERS`. Availability does not establish candidate analysis. See the [runtime guide](mi355x-runtime.md#sdk-and-profiling). | | Docker | Current stable release | Required; serial experiments run through `make docker-run`; multi-GPU experiments run through `make docker-parallel-run`. | -| SGLang runtime image | `lmsysorg/sglang:v0.5.12-rocm720-mi30x` for `gfx942`; `lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705` for `gfx950` | The verified `gfx950` digest is `sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78`. Override with `AKA_DOCKER_IMAGE`, `AKA_DOCKER_IMAGE_GFX942`, or `AKA_DOCKER_IMAGE_GFX950`. | +| SGLang runtime image | `lmsysorg/sglang:v0.5.12-rocm720-mi30x` for `gfx942`; SGLang 0.5.20 / ROCm 10 for `gfx950` | The MI355X default is `lmsysorg/sglang@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69`. Override with `AKA_DOCKER_IMAGE`, `AKA_DOCKER_IMAGE_GFX942`, or `AKA_DOCKER_IMAGE_GFX950`. Task-specific runtime requirements still apply. | | RDNA4 runtime image | [Digest-pinned base and layout adapter](../../docker/rdna4/Dockerfile) | Default image builds on first use if missing; `make docker-build-rdna4` prebuilds or rebuilds it. Image overrides disable automatic builds; see the [runtime guide](../../docker/rdna4/README.md). | | Python for GPU experiments | Provided by the qualified task runtime image | Task-specific qualification applies; older images do not support every retained task. | | Python for the full repository CPU/source audit | CPython 3.12 | The complete source set requires 3.12 syntax, and migration AST fingerprints are pinned to this minor version. Use full Git history; see [contributor verification](../../CONTRIBUTING.md#testing-and-verification). | | Node.js and npm | Node.js 22+ for Claude Code's npm installation; Node.js 24 for DeepSeek Harness | Required on the host for npm-installed agent CLIs; DeepSeek uses a dedicated prefix. | | PyTorch | ROCm build bundled in the image | Provided by the selected runtime image. | | Triton | Bundled with the image's ROCm PyTorch | Required for Triton task categories. | -| AITER | `0.1.17.dev110+g9127c94a1` in the verified `gfx950` image | Required by AITER-backed task oracles and kernels. | -| FlyDSL | `0.2.2` in the verified `gfx950` image (or `make docker-setup-flydsl` when absent) | Required for `flydsl2flydsl`, `torch2flydsl`, `triton2flydsl`, and `operator2flydsl` tasks. | +| AITER | Source installation bundled in the pinned `gfx950` image | Required by AITER-backed task oracles and kernels. Distribution metadata alone may not identify a source installation; retain the image identity and materialized source hashes. | +| FlyDSL | `0.3.2` in the default `gfx950` image (or `make docker-setup-flydsl` when absent) | Required for `flydsl2flydsl`, `torch2flydsl`, `triton2flydsl`, and `operator2flydsl` tasks. | ## Evaluation-tool sidecars Optional Triton FpSan, GPU ASan, rocJITsu Race Detector, rocJITsu Waitcheck, rocJITsu ConSan, and HIP-FpSan dependencies are kept out of the scoring image -and installed in one isolated sidecar image per tool. The scoring image, -FlyDSL, and AITER versions in the preceding table remain unchanged. +and installed in one isolated sidecar image per tool. Enabling a tool does not +install packages into the scoring image or change its dependency versions. | GPU architecture | Sidecar status | Notes | | --- | --- | --- | -| `gfx950` (MI355X) | Runtime-qualified, candidate-dependent | Pinned image/build locks and all six integrated startup controls pass on the current hardware. End-to-end readiness still depends on language, artifact, adapter, and candidate attestation. Waitcheck and ConSan are qualified only for explicitly configured advisory pilots. Trusted single-dispatch Triton/FlyDSL rocJITsu capsule replay is implemented, but automatic evaluator-owned capsule capture and binding to the correctness run remain advisory-only gaps. | +| `gfx950` (MI355X) | Migration qualification incomplete, candidate-dependent | Five ROCm 10 runtime startup controls pass; GPU ASan safe probes fail on the tested hosts and remain unqualified. See the [runtime guide](mi355x-runtime.md). End-to-end readiness still depends on language, artifact, adapter, and candidate attestation. Waitcheck and ConSan are qualified only for explicitly configured advisory pilots. Trusted single-dispatch Triton/FlyDSL rocJITsu capsule replay is implemented, but automatic evaluator-owned capsule capture and binding to the correctness run remain advisory-only gaps. | | `gfx942` (MI300X/MI325X) | Unverified | No equivalent image/adapter/positive-control qualification has completed; the host runner currently rejects evaluation-tool sidecars. | | `gfx1201` (RDNA4) | Unverified | Runtime task checks do not qualify evaluation-tool sidecars; the host runner rejects them. | @@ -99,9 +99,7 @@ validated on hardware. Model/provider support and run-level overrides are integration-specific; there is no shared top-level provider field. -| Provider | Notes | -| --- | --- | -| OpenAI | Use a selected integration or CLI configured for OpenAI. | -| Anthropic | Use a selected integration or CLI configured for Anthropic. | -| DeepSeek | `deepseek_harness` selects the upstream DeepSeek provider; model, protocol, and optional endpoint override are integration-local settings. | -| OpenRouter or another OpenAI-compatible service | Supported when the selected integration accepts a custom provider/base URL. | +For supported providers, model selection and custom endpoint settings, use the +[selected integration's configuration guide](../how-to/agents.md) and its +`agents//agent_config.yaml`. Provider configuration remains local +to that integration. diff --git a/docs/reference/mi355x-runtime.md b/docs/reference/mi355x-runtime.md new file mode 100644 index 000000000..a64d00aa4 --- /dev/null +++ b/docs/reference/mi355x-runtime.md @@ -0,0 +1,62 @@ +# MI355X runtime and compatibility + +The proposed MI355X default is the immutable SGLang 0.5.20 / ROCm 10 manifest +selected in [`docker_benchmark.sh`](../../src/scripts/docker_benchmark.sh). +The `gfx942` and `gfx1201` defaults are unchanged. Explicit image overrides and +task-specific runtime requirements still apply. The migration remains +**unqualified for promotion** while GPU ASan cannot pass its required controls. +Earlier image qualifications do not qualify this source/image combination. + +All six optional tool runtimes use the same scoring base and distinct +`gfx950-rocm10` tags. The scoring-image verifier rejects mismatched images and +freezes accepted aliases to their local image IDs. See the maintained +[tool capability matrix](../how-to/use-evaluation-tools.md) for startup and +candidate-analysis boundaries. GPU ASan needs working GPU access to ordinary +host allocations; a successful build or an `HSA_XNACK=1` setting alone does not +establish that prerequisite. Required safe and known-bug probes remain mandatory. + +Ordinary runs now propagate their selected image into materialization and session +identity. Start a fresh run after changing the image, including for workspaces +created before image identity was recorded. Custom mutable tags must be frozen +to a digest. Keep earlier artifacts and compare baseline and candidate within +the same runtime; do not treat an image change as a kernel optimization. + +The new runtime contains Python 3.12, PyTorch 2.11, Triton 3.8 and FlyDSL 0.3.2. +ROCm SDK and HIP component versions are distinct. AITER is installed from source +without distribution metadata, so retain the image identity and materialized +source hashes. Image-acquired AITER trees can change even when task config and +candidate bytes do not. Tasks that require a vLLM installation still need their +explicitly documented runtime; the SGLang default does not provide it. + +## SDK and profiling + +The runner selects SDK core libraries before starting workloads and routes +`rocprofv3` through that same tree. Loading both the SDK developer and core copies +of COMGR can abort a Python workload with duplicate LLVM registration, including +when vLLM loads a PyTorch library before importing PyTorch. `rocprof-compute` is +optional and absent from the base image; profiler availability is separate from +successful trace/counter collection and from ordinary graph/event timing. +GEAK's SDK cache is separated by Python ABI to prevent reuse of Python 3.10 +extension modules in Python 3.12. + +## Tasks and analysis tools + +FlyDSL tasks use task-local compatibility adapters for removed buffer/vector +APIs and retain their existing cases, numerical gates and timing contracts. +The gfx950 MLA decode path uses a single-stage schedule because its two-stage +pipeline produced incorrect results with the new compiler. This fallback can +increase latency. Numerical failures in the FP8 fused-MoE task have also been +reproduced on the previous runtime and remain unresolved; they must not be +counted as passing qualification or attributed to this upgrade without evidence. + +Triton AOT extraction resolves named signature keys against function argument +names and rejects unrepresented global scratch. FlyDSL extraction accepts the +asynchronous launch form while rejecting additional launches. GPU ASan candidate +invocations use the same attested HIP/HSA and library paths as startup probes; +those paths cannot be supplied or overridden by task configuration. + +Keep current validation reports, source/config fingerprints, scheduler outcomes +and old/new comparison data outside the repository. Report completed checks and +remaining limitations in the PR. Each changed task needs a fresh full +framework-finalized validator PASS; initial actions, CPU tests and interrupted +runs do not replace that requirement. diff --git a/docs/reference/release-notes.md b/docs/reference/release-notes.md index c974f371b..0a41c5c33 100644 --- a/docs/reference/release-notes.md +++ b/docs/reference/release-notes.md @@ -141,8 +141,9 @@ The task validator now includes Codex backend support, improved Python-environme - Evaluation-tool sidecars are experimental and `gfx950`-only. Startup controls prove a tool installation can detect its synthetic bug, not that a candidate was instrumented. No bundled task currently supplies a production-qualified - adapter/attestation. All six integrated startup controls pass on the current - MI355X qualification host. Synthetic manager-to-sidecar candidate pairs also + adapter/attestation. All six integrated startup controls passed on the + ROCm 7.2 MI355X qualification host; the ROCm 10 profile has separate + [qualification requirements](runtime-upgrade-qualification.md). Synthetic manager-to-sidecar candidate pairs also distinguished clean from seeded-bug Triton FpSan, HIP/Triton GPU ASan, and HIP-FpSan runs; trusted AOT replay produced a clean Triton result and found the seeded FlyDSL LDS race. Waitcheck distinguished a correct wait from a missing diff --git a/docs/sphinx/_toc.yml.in b/docs/sphinx/_toc.yml.in index ef12e4729..543ab652c 100644 --- a/docs/sphinx/_toc.yml.in +++ b/docs/sphinx/_toc.yml.in @@ -57,6 +57,8 @@ subtrees: title: Configuration and API reference - file: reference/benchmark-methodology title: Performance measurement methodology + - file: reference/mi355x-runtime + title: MI355X runtime and compatibility - url: https://rocm.docs.amd.com/projects/hyperloom/en/latest/index.html title: Hyperloom diff --git a/example_configs/evaluation_tools_advisory_mi355x.yaml b/example_configs/evaluation_tools_advisory_mi355x.yaml index 66416cef9..ad1ffd0ec 100644 --- a/example_configs/evaluation_tools_advisory_mi355x.yaml +++ b/example_configs/evaluation_tools_advisory_mi355x.yaml @@ -27,7 +27,7 @@ workspace_directory_prefix: workspace evaluation_tools: # Replace false with a list such as [gpu_asan, rocjitsu] only after the - # selected tasks define matching adapters. `true` enables all four tools and + # selected tasks define matching adapters. `true` enables all six tools and # is rarely appropriate for a heterogeneous run. Host AKA_EVAL_TOOLS, when # set, authoritatively replaces this subset for both sidecar startup and plan. enabled: false @@ -39,7 +39,7 @@ evaluation_tools: # AKA_EVAL_TOOL_IMAGE_ variable selects each image; the runner resolves # and injects its immutable local image ID into the plan automatically. # The scoring image must also resolve to the same immutable image ID as the - # pinned SGLang gfx950 content-addressed manifest reference; its selected + # pinned ROCm 10 SGLang gfx950 content-addressed manifest reference; its selected # reference and ID are added to plan source evidence. # Do not place positive_control_required or any runtime binary/library/header # path in options; those keys are reserved and configuration loading rejects diff --git a/src/eval_tools/README.md b/src/eval_tools/README.md index b580bbe6e..4c44cf9e5 100644 --- a/src/eval_tools/README.md +++ b/src/eval_tools/README.md @@ -25,7 +25,7 @@ would complement the current rocJITsu checks with compiler-instrumented dynamic race detection on applicable device code. TSAN is not currently registered as an AgentKernelArena evaluation tool. The -pinned ROCm 7.2 HIP compiler used by the `gfx950` scoring baseline warns that +previous ROCm 7.2 baseline compiler warned that `-fsanitize=thread` is unsupported for the `amdgcn-amd-amdhsa` device target and ignores the option. Host `libclang_rt.tsan` files do not establish device instrumentation or GPU runtime coverage. Public ROCm development documentation @@ -33,6 +33,9 @@ describes device-side TSAN builds for `gfx942` and `gfx950`, but that project development path has not been qualified as a pinned, general-purpose evaluator runtime for this benchmark. +The ROCm 10 default-image migration does not add a TSAN runtime or qualify +device TSAN coverage. + Future TSAN support should be added only after a compatible device compiler and runtime can be pinned in an isolated sidecar. Promotion requires exact candidate recompilation and build attestation, safe/racy startup controls, end-to-end @@ -53,11 +56,14 @@ host launchers and native support code, but host instrumentation does not prove that an optimized GPU kernel was checked. UBSAN is not currently registered as an AgentKernelArena evaluation tool. ROCm -documents its current development build as host-only, and the pinned ROCm 7.2 -HIP compiler warns that `-fsanitize=undefined` is unsupported for the +development documentation describes host-only support, and the previous ROCm 7.2 +HIP compiler warned that `-fsanitize=undefined` is unsupported for the `amdgcn-amd-amdhsa` device target and ignores the option. AgentKernelArena must therefore not treat a host-only UBSAN run as device-kernel sanitizer coverage. +The ROCm 10 default-image migration does not add a UBSAN runtime or qualify +device UBSAN coverage. + Future GPU UBSAN support depends on upstream device instrumentation and a compatible device runtime becoming available. Once available, it will require the same isolated sidecar, exact-artifact attestation, positive/negative diff --git a/src/eval_tools/adapters/flydsl_aot.py b/src/eval_tools/adapters/flydsl_aot.py index 122a87902..92aaf6ed0 100644 --- a/src/eval_tools/adapters/flydsl_aot.py +++ b/src/eval_tools/adapters/flydsl_aot.py @@ -13,7 +13,11 @@ _CONST_RE = re.compile(r"(%[A-Za-z0-9_.$-]+)\s*=\s*arith\.constant\s+(-?\d+)\s*:\s*(?:index|i\d+)") -_LAUNCH_RE = re.compile(r"gpu\.launch_func\s+@(?:[A-Za-z0-9_.$-]+::)?@(?P[A-Za-z0-9_.$-]+)(?P.*?)(?=\n\s*[%}]|$)", re.S) +_LAUNCH_RE = re.compile( + r"gpu\.launch_func\s+(?:async\s*\[[^\]]*\]\s*)?" + r"@(?:[A-Za-z0-9_.$-]+::)?@(?P[A-Za-z0-9_.$-]+)" + r"(?P.*?)(?=\n\s*[%}]|$)", re.S, +) _BIN_RE = re.compile(r'\bbin\s*=\s*"((?:\\.|[^"\\])*)"', re.S) @@ -57,7 +61,7 @@ def _parse_dims(body: str, label: str, constants: Mapping[str, int]) -> tuple[in def parse_flydsl_static_launch(source_ir: str) -> FlyDslStaticLaunch: constants = {name: int(value) for name, value in _CONST_RE.findall(source_ir)} launches = list(_LAUNCH_RE.finditer(source_ir)) - if len(launches) != 1: + if len(launches) != 1 or len(re.findall(r"\bgpu\.launch_func\b", source_ir)) != 1: raise CapsuleValidationError(f"FlyDSL replay MVP requires one gpu.launch_func, found {len(launches)}") match = launches[0] body = match.group("body") diff --git a/src/eval_tools/adapters/triton_aot.py b/src/eval_tools/adapters/triton_aot.py index 5ee1cbb34..a5f23ce8d 100644 --- a/src/eval_tools/adapters/triton_aot.py +++ b/src/eval_tools/adapters/triton_aot.py @@ -114,6 +114,40 @@ def extract_triton_aot( if not isinstance(signature, Mapping): raise CapsuleValidationError("CompiledKernel source signature is unavailable") + # ASTSource uses parameter names in newer Triton releases; IRSource and + # older releases use indices. Resolve names from the function declaration, + # never mapping insertion order (which need not match the launcher ABI). + arg_names = getattr(getattr(src, "fn", None), "arg_names", ()) + if not isinstance(arg_names, (tuple, list)) or any(not isinstance(n, str) for n in arg_names): + arg_names = () + if len(set(arg_names)) != len(arg_names): + raise CapsuleValidationError("Triton argument names are ambiguous") + + def argument_index(key: Any) -> int: + if isinstance(key, str) and key in arg_names: + return arg_names.index(key) + if isinstance(key, int) and not isinstance(key, bool) and key >= 0: + return key + if isinstance(key, str) and key.isdecimal(): + return int(key) + raise CapsuleValidationError(f"Triton argument {key!r} has no source position") + + indexed_signature: dict[int, str] = {} + for key, value in signature.items(): + index = argument_index(key) + if index in indexed_signature: + raise CapsuleValidationError(f"duplicate Triton argument position {index}") + indexed_signature[index] = str(value) + indexed_constants = {} + if not isinstance(constants, Mapping): + raise CapsuleValidationError("CompiledKernel source constants are unavailable") + for key, value in constants.items(): + if isinstance(key, tuple): + if len(key) != 1: + raise CapsuleValidationError("nested Triton constexpr arguments are not supported") + key = key[0] + indexed_constants[argument_index(key)] = value + dims = tuple(int(v) for v in grid) if len(dims) > 3 or not dims: raise CapsuleValidationError("Triton grid must have one to three dimensions") @@ -125,6 +159,8 @@ def extract_triton_aot( launch = LaunchSpec(grid3, (num_warps * warp_size, 1, 1), int(_metadata_value(metadata, "shared", 0))) launch.validate() profile_per_grid = int(_metadata_value(metadata, "profile_scratch_size", 0)) + if int(_metadata_value(metadata, "global_scratch_size", 0)) != 0: + raise CapsuleValidationError("Triton global scratch allocation is not supported by the AOT replay MVP") scratch = ScratchSpec( global_bytes=0, profile_bytes=profile_per_grid * grid3[0] * grid3[1] * grid3[2], @@ -132,8 +168,8 @@ def extract_triton_aot( ) scratch.validate() abi = build_triton_abi( - {int(k): str(v) for k, v in signature.items()}, - constants=constants if isinstance(constants, Mapping) else {}, + indexed_signature, + constants=indexed_constants, pointer_bindings=pointer_bindings, scalar_values=scalar_values, ) diff --git a/src/eval_tools/config.py b/src/eval_tools/config.py index 833be787c..ec8d9e4ed 100644 --- a/src/eval_tools/config.py +++ b/src/eval_tools/config.py @@ -35,6 +35,8 @@ { "asan_runtime_dir", "hip_asan_runtime", + "hsa_asan_runtime", + "asan_extra_library_dirs", "host_asan_preload", "host_asan_lib_dir", "normal_rocm_lib_dir", diff --git a/src/eval_tools/contracts.py b/src/eval_tools/contracts.py index 2006c5364..3764b5088 100644 --- a/src/eval_tools/contracts.py +++ b/src/eval_tools/contracts.py @@ -16,7 +16,7 @@ class _StringEnum(str, Enum): - """``StrEnum`` compatible with the Python 3.10 scoring image.""" + """``StrEnum`` compatible with supported Python 3.10+ runtimes.""" def __str__(self) -> str: return self.value diff --git a/src/eval_tools/plugins/gpu_asan.py b/src/eval_tools/plugins/gpu_asan.py index 931983977..4d53a04b9 100644 --- a/src/eval_tools/plugins/gpu_asan.py +++ b/src/eval_tools/plugins/gpu_asan.py @@ -73,7 +73,7 @@ def assess(self, context: ToolContext, runtime: CapabilityCheck) -> ToolCapabili engine = blocked_check( CapabilityState.UNSUPPORTED, "gpu_asan_flydsl_no_device_instrumentation", - "FlyDSL 0.2.x does not insert the AMDGPU AddressSanitizer pass.", + "No qualified FlyDSL GPU ASan instrumentation adapter is available.", ) elif profile.framework == "aiter" or profile.artifact_kind == ArtifactKind.HSACO_PRECOMPILED: # Explicit source-rebuild evidence may make a future AITER HIP lane @@ -148,26 +148,35 @@ def build_invocation(self, context: ToolContext) -> ToolInvocation: # probe and intentionally never stat'ed in the scoring container. host_runtime = sidecar_path(context, "host_asan_preload") hip_runtime = sidecar_path(context, "hip_asan_runtime") + hsa_runtime = sidecar_path(context, "hsa_asan_runtime") runtime_dir = sidecar_path(context, "asan_runtime_dir") host_library_dir = sidecar_path(context, "host_asan_lib_dir") normal_rocm_library_dir = sidecar_path(context, "normal_rocm_lib_dir") + extra_library_dirs = context.options.get("asan_extra_library_dirs", ()) + if not isinstance(extra_library_dirs, (list, tuple)) or any( + not isinstance(path, str) or "\x00" in path or not Path(path).is_absolute() + for path in extra_library_dirs + ): + raise ValueError("asan_extra_library_dirs must be a list of absolute sidecar paths") if host_runtime is not None and host_library_dir is None: host_library_dir = host_runtime.parent env["LD_LIBRARY_PATH"] = _prepend_paths( str(host_library_dir) if host_library_dir else None, str(runtime_dir) if runtime_dir else None, + *extra_library_dirs, str(normal_rocm_library_dir) if normal_rocm_library_dir else None, inherited=env.get("LD_LIBRARY_PATH", ""), ) - if is_triton: - if host_runtime is None or hip_runtime is None: - raise ValueError( - "Triton GPU ASan requires attested host_asan_preload and " - "hip_asan_runtime sidecar paths" - ) + if is_triton and (host_runtime is None or hip_runtime is None): + raise ValueError( + "Triton GPU ASan requires attested host_asan_preload and " + "hip_asan_runtime sidecar paths" + ) + if host_runtime is not None and hip_runtime is not None: env["LD_PRELOAD"] = _prepend_paths( str(host_runtime), str(hip_runtime), + str(hsa_runtime) if hsa_runtime else None, inherited=env.get("LD_PRELOAD", ""), ) attestation_path = artifact_path( diff --git a/src/eval_tools/runtime_client.py b/src/eval_tools/runtime_client.py index 71332f4a1..5390a48c4 100644 --- a/src/eval_tools/runtime_client.py +++ b/src/eval_tools/runtime_client.py @@ -377,6 +377,7 @@ def probe(self, tool: str, context: ToolContext) -> CapabilityCheck: "gpu_asan": ( "asan_runtime_dir", "hip_asan_runtime", + "hsa_asan_runtime", "host_asan_preload", "host_asan_lib_dir", "normal_rocm_lib_dir", diff --git a/src/eval_tools/worker.py b/src/eval_tools/worker.py index 3ed7fb35d..5ec994046 100644 --- a/src/eval_tools/worker.py +++ b/src/eval_tools/worker.py @@ -36,7 +36,7 @@ MAX_ARGV_ENTRIES = 4096 MAX_ENV_ENTRIES = 4096 MAX_LOG_LIMIT_BYTES = 1024 * 1024 * 1024 -TRITON_FPSAN_VERSION = "3.7.0+amd.rocm7.2.0.gitd0d77a509" +TRITON_FPSAN_VERSION = "3.8.0+git4cff872c.rocm10.0.0" TRITON_ASAN_VERSIONS = { "3.6.0+git42270451", TRITON_FPSAN_VERSION, @@ -313,10 +313,37 @@ def _positive_result( def _triton_fpsan_positive( probe_root: Path, work_dir: Path, artifact_dir: Path ) -> dict[str, Any]: + steps = {} + controls = {} + for name, arguments, expected_equal in ( + ("equivalent", [], True), + ("known_mismatch", ["wrong"], False), + ): + result = _triton_fpsan_control( + probe_root, work_dir / name, artifact_dir / name, + arguments=arguments, expected_equal=expected_equal, + ) + steps.update({f"{name}_{key}": value for key, value in result["steps"].items()}) + controls[name] = {"passed": result["passed"]} + return _positive_result( + passed=all(control["passed"] for control in controls.values()), + kind="triton_fpsan_equivalence_and_mismatch", + detail="equivalent and known-wrong expressions must have the expected digest relation with independently attested FpSan compilation", + artifact_dir=artifact_dir, + steps=steps, + controls=controls, + ) + + +def _triton_fpsan_control( + probe_root: Path, work_dir: Path, artifact_dir: Path, + *, arguments: list[str], expected_equal: bool, +) -> dict[str, Any]: + work_dir.mkdir(parents=True, exist_ok=True) cache = work_dir / "triton-fpsan-cache" step = _run_probe_step( - "known-mismatch", - [shutil.which("python") or os.sys.executable, str(probe_root / "triton_fpsan_probe.py"), "wrong"], + "comparison", + [shutil.which("python") or os.sys.executable, str(probe_root / "triton_fpsan_probe.py"), *arguments], cwd=work_dir, environment={ "TRITON_INSTRUMENTATION_MODE": "fpsan", @@ -337,44 +364,46 @@ def _triton_fpsan_positive( continue if isinstance(metadata, dict): modes.append(metadata.get("instrumentation_mode")) - mismatch = bool( + expected_relation = bool( record and record.get("reference_digest") and record.get("candidate_digest") - and record["reference_digest"] != record["candidate_digest"] + and (record["reference_digest"] == record["candidate_digest"]) == expected_equal ) - passed = step["returncode"] == 0 and mismatch and len(modes) >= 2 and all( + passed = step["returncode"] == 0 and expected_relation and len(modes) >= 2 and all( mode == "fpsan" for mode in modes ) return _positive_result( passed=passed, - kind="triton_fpsan_known_mismatch", + kind="triton_fpsan_comparison", detail=( - "known numerical mismatch produced distinct digests and every compiled kernel metadata record attested instrumentation_mode=fpsan" + "expected digest relation and every compiled kernel metadata record attested instrumentation_mode=fpsan" if passed - else "known mismatch did not produce both distinct digests and compiler metadata attesting fpsan mode" + else "comparison did not produce both the expected digest relation and compiler metadata attesting fpsan mode" ), artifact_dir=artifact_dir, - steps={"known_mismatch": step}, + steps={"comparison": step}, ) def _asan_environment(evidence: Mapping[str, Any]) -> dict[str, str]: host_preload = str(evidence.get("host_asan_preload") or "") hip_runtime = str(evidence.get("hip_asan_runtime") or "") + hsa_runtime = str(evidence.get("hsa_asan_runtime") or "") host_lib_dir = str(evidence.get("host_asan_lib_dir") or "") runtime_dir = str(evidence.get("asan_runtime_dir") or "") normal_rocm_lib = str(evidence.get("normal_rocm_lib_dir") or "") inherited_preload = os.environ.get("LD_PRELOAD", "") inherited_libraries = os.environ.get("LD_LIBRARY_PATH", "") preload = ":".join( - value for value in (host_preload, hip_runtime, inherited_preload) if value + value for value in (host_preload, hip_runtime, hsa_runtime, inherited_preload) if value ) library_path = ":".join( value for value in ( host_lib_dir, runtime_dir, + *evidence.get("asan_extra_library_dirs", ()), normal_rocm_lib, inherited_libraries, ) @@ -981,7 +1010,7 @@ def runtime_evidence( "expected_triton_version": TRITON_FPSAN_VERSION, }) elif tool == "gpu_asan": - runtime_dir = Path(os.environ.get("AKA_GPU_ASAN_RUNTIME_DIR", "/opt/rocm-7.2.0/lib/asan")) + runtime_dir = Path(os.environ.get("AKA_GPU_ASAN_RUNTIME_DIR", "/opt/rocm-asan-10.0/lib")) hip_runtime = Path( os.environ.get( "AKA_GPU_ASAN_HIP_RUNTIME", str(runtime_dir / "libamdhip64.so") @@ -993,7 +1022,7 @@ def runtime_evidence( runtime_dir / "libamd_comgr.so", ) preload_candidates = sorted( - Path("/opt/rocm-7.2.0/lib/llvm/lib/clang").glob( + (runtime_dir / "llvm/lib/clang").glob( "*/lib/linux/libclang_rt.asan-x86_64.so" ) ) @@ -1007,9 +1036,17 @@ def runtime_evidence( evidence.update({ "asan_runtime_dir": str(runtime_dir), "hip_asan_runtime": str(hip_runtime) if hip_runtime.is_file() else None, + "hsa_asan_runtime": str(runtime_dir / "libhsa-runtime64.so") + if (runtime_dir / "libhsa-runtime64.so").is_file() else None, + "asan_extra_library_dirs": [ + str(path) for path in (runtime_dir / "llvm/lib", runtime_dir / "rocm_sysdeps/lib") + if path.is_dir() + ], "host_asan_preload": str(host_preload) if host_preload and host_preload.is_file() else None, "host_asan_lib_dir": str(host_preload.parent) if host_preload and host_preload.is_file() else None, - "normal_rocm_lib_dir": "/opt/rocm-7.2.0/lib", + "normal_rocm_lib_dir": os.environ.get( + "AKA_GPU_ASAN_NORMAL_ROCM_LIB_DIR", "/opt/rocm/lib" + ), "asan_libraries": {path.name: path.is_file() for path in required_libraries}, "triton_version": triton_version, "triton_asan": triton_version in TRITON_ASAN_VERSIONS, diff --git a/src/scripts/docker_benchmark.sh b/src/scripts/docker_benchmark.sh index 7051bc6ba..cf88acaba 100755 --- a/src/scripts/docker_benchmark.sh +++ b/src/scripts/docker_benchmark.sh @@ -9,9 +9,8 @@ GFX950_V0519_DOCKER_IMAGE="lmsysorg/sglang-rocm:v0.5.19-rocm10-mi35x-20260913" GFX950_V0519_IMMUTABLE_IMAGE="lmsysorg/sglang-rocm@sha256:106a7adbeec5554b6e66a4bda0b3694af442717b9fe92754a9885520077b6f93" GFX950_V0520_DOCKER_IMAGE="lmsysorg/sglang:v0.5.20-rocm10-mi35x" GFX950_V0520_IMMUTABLE_IMAGE="lmsysorg/sglang@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69" -# Keep qualification and scoring on the verified bytes, even if the dated tag -# moves. New runtime candidates remain explicit overrides until qualified. -DEFAULT_DOCKER_IMAGE_GFX950="${AKA_DOCKER_IMAGE_GFX950:-$GFX950_V0514_IMMUTABLE_IMAGE}" +# Pin the MI355X runtime; retain the older manifest for explicit rollback. +DEFAULT_DOCKER_IMAGE_GFX950="${AKA_DOCKER_IMAGE_GFX950:-$GFX950_V0520_IMMUTABLE_IMAGE}" # Built on first use when absent; its Dockerfile pins the base. DEFAULT_DOCKER_IMAGE_GFX1201="${AKA_DOCKER_IMAGE_GFX1201:-agent-kernel-arena:rdna4-rocm10-v1}" CONTAINER_WORKDIR="${AKA_DOCKER_WORKDIR:-/workspace}" @@ -42,6 +41,7 @@ EVAL_TOOL_FRAMEWORK_CONTAINER_ROOT="/opt/aka-eval-tools" EVAL_TOOL_SCRATCH_CONTAINER_DIR="/work" EVAL_TOOL_ARTIFACT_CONTAINER_DIR="/artifacts" EVAL_TOOL_RUNTIME_DIR="" +EVAL_TOOL_SCORING_SDK_RUNTIME=0 EVAL_TOOL_SOCKET_HOST_DIR="" EVAL_TOOL_ARTIFACT_HOST_ROOT="" EVAL_TOOL_ARTIFACT_SCORING_ROOT="" @@ -156,7 +156,10 @@ uses_gfx950_aiter_cache_overrides() { [[ "$SELECTED_GPU_ARCH" == "gfx950" ]] || return 1 # Evaluation-tool setup may replace SELECTED_IMAGE with its verified local # ID. That verifier checks the pinned scoring bytes, including custom aliases. - local image_reference="${AKA_SCORING_IMAGE_REFERENCE:-$SELECTED_IMAGE}" + local image_reference="$SELECTED_IMAGE" + if [[ -n "${AKA_EVAL_TOOL_SOCKET_HOST_DIR:-}" ]]; then + image_reference="${AKA_SCORING_IMAGE_REFERENCE:-$SELECTED_IMAGE}" + fi [[ "$image_reference" == "$GFX950_V0514_DOCKER_IMAGE" \ || "$image_reference" == "$GFX950_V0514_IMMUTABLE_IMAGE" \ || "$image_reference" == "$GFX950_V0519_DOCKER_IMAGE" \ @@ -707,7 +710,7 @@ eval_tool_image() { if [[ -n "$override" ]]; then printf '%s\n' "$override" else - printf 'agent-kernel-arena/eval-tool-%s:gfx950\n' "${tool//_/-}" + printf 'agent-kernel-arena/eval-tool-%s:gfx950-rocm10\n' "${tool//_/-}" fi } @@ -716,15 +719,16 @@ verify_eval_tool_scoring_image() { [[ "$SELECTED_GPU_ARCH" == "gfx950" ]] \ || die "Evaluation-tool sidecars are verified only for gfx950" selected_id="$(docker image inspect --format '{{.Id}}' "$SELECTED_IMAGE" 2>/dev/null || true)" - pinned_id="$(docker image inspect --format '{{.Id}}' "$GFX950_V0514_IMMUTABLE_IMAGE" 2>/dev/null || true)" + pinned_id="$(docker image inspect --format '{{.Id}}' "$GFX950_V0520_IMMUTABLE_IMAGE" 2>/dev/null || true)" [[ "$selected_id" == sha256:* ]] \ || die "Could not resolve immutable scoring image ID for $SELECTED_IMAGE" [[ "$pinned_id" == sha256:* ]] \ - || die "Pinned evaluation scoring image is unavailable: $GFX950_V0514_IMMUTABLE_IMAGE" + || die "Pinned evaluation scoring image is unavailable: $GFX950_V0520_IMMUTABLE_IMAGE" [[ "$selected_id" == "$pinned_id" ]] \ - || die "Evaluation tools are unverified with scoring image $SELECTED_IMAGE ($selected_id); expected $GFX950_V0514_DOCKER_IMAGE ($pinned_id)" + || die "Evaluation tools are unverified with scoring image $SELECTED_IMAGE ($selected_id); expected $GFX950_V0520_DOCKER_IMAGE ($pinned_id)" export AKA_SCORING_IMAGE_RUNTIME_REF="$selected_id" export AKA_SCORING_IMAGE_REFERENCE="$SELECTED_IMAGE" + EVAL_TOOL_SCORING_SDK_RUNTIME=1 # Launch by immutable local config ID after verification. This closes the # gap in which a mutable tag could move between inspection and docker run. SELECTED_IMAGE="$selected_id" @@ -1057,6 +1061,10 @@ build_docker_args() { # Parallel runs finish their preflight before starting any workers, so the # first container builds a missing default and subsequent containers reuse it. ensure_runtime_image + local scoring_image_reference="$SELECTED_IMAGE" + if [[ -n "${AKA_EVAL_TOOL_SOCKET_HOST_DIR:-}" ]]; then + scoring_image_reference="${AKA_SCORING_IMAGE_REFERENCE:?missing verified scoring image reference}" + fi docker_args=(run --rm --entrypoint bash) unset _MOUNTED_TARGETS @@ -1090,6 +1098,9 @@ build_docker_args() { -e "AGENT_KERNEL_ARENA_DOCKER=1" -e "AGENT_KERNEL_ARENA_WORKDIR=${CONTAINER_WORKDIR}" -e "AGENT_KERNEL_ARENA_GPU_ARCH=${SELECTED_GPU_ARCH}" + # Bind ordinary runs as well as tool-enabled runs to their selected + # runtime. Materialization and session resume compare this identity. + -e "AKA_SCORING_IMAGE_REFERENCE=$scoring_image_reference" -e "AKA_REQUIRED_PROFILERS=${AKA_REQUIRED_PROFILERS:-}" -e "PYTORCH_ROCM_ARCH=${SELECTED_GPU_ARCH}" -e "AGENT_STATE_MOUNT_ROOT=${AGENT_STATE_MOUNT_ROOT}" @@ -1097,9 +1108,20 @@ build_docker_args() { -w "$CONTAINER_WORKDIR" ) + local sdk_core_runtime="$EVAL_TOOL_SCORING_SDK_RUNTIME" + case "$scoring_image_reference" in + "$GFX950_V0520_DOCKER_IMAGE"|"$GFX950_V0520_IMMUTABLE_IMAGE"|"${GFX950_V0520_DOCKER_IMAGE}@${GFX950_V0520_IMMUTABLE_IMAGE##*@}") + sdk_core_runtime=1 + ;; + esac + if [[ "$sdk_core_runtime" == "1" ]]; then + docker_args+=(-e "AKA_ROCM_SDK_CORE_RUNTIME=1") + fi + # GEAK's claude-agent-sdk is installed with `pip install --target` into # this host-mounted dir (see container_setup_geak). Only put it on - # a GEAK-only path for the container bootstrap to prepend to PYTHONPATH. + # a GEAK-only root for the bootstrap to qualify by Python ABI and prepend + # to PYTHONPATH. A ROCm image upgrade can change that ABI. # Do not replace the image's PYTHONPATH: it can supply AITER/source imports. if [[ "$GEAK_RUNTIME" == "1" ]]; then docker_args+=(-e "AKA_GEAK_SDK_PATH=${CONTAINER_WORKDIR}/.aka-pyuserbase/geak-sdk") @@ -1265,7 +1287,6 @@ build_docker_args() { -e "AKA_EVAL_TOOL_SCORING_ROOT=$CONTAINER_WORKDIR" -e "AKA_EVAL_TOOL_ARTIFACT_SCORING_ROOT=$AKA_EVAL_TOOL_ARTIFACT_SCORING_ROOT" -e "AKA_SCORING_IMAGE_RUNTIME_REF=${AKA_SCORING_IMAGE_RUNTIME_REF:?missing verified scoring image ID}" - -e "AKA_SCORING_IMAGE_REFERENCE=${AKA_SCORING_IMAGE_REFERENCE:?missing scoring image reference}" ) if [[ -n "${AKA_EVAL_TOOLS_SELECTED:-}" ]]; then docker_args+=(-e "AKA_EVAL_TOOLS_SELECTED=${AKA_EVAL_TOOLS_SELECTED}") @@ -1327,7 +1348,7 @@ docker_exec() { local interactive="${1:-0}" shift build_docker_args "$interactive" - docker "${docker_args[@]}" -lc 'cd "$AGENT_KERNEL_ARENA_WORKDIR" && if [[ -n "${AKA_GEAK_SDK_PATH:-}" ]]; then export PYTHONPATH="${AKA_GEAK_SDK_PATH}${PYTHONPATH:+:$PYTHONPATH}"; fi && if [[ "${AGENT_KERNEL_ARENA_ISOLATED_HOME:-0}" == "1" ]]; then bash src/scripts/docker_benchmark.sh _container_prepare_worker_home; fi && exec "$@"' _ "$@" + docker "${docker_args[@]}" -lc 'cd "$AGENT_KERNEL_ARENA_WORKDIR" && if [[ -n "${AKA_GEAK_SDK_PATH:-}" ]]; then export AKA_GEAK_SDK_PATH="${AKA_GEAK_SDK_PATH}/$(python -c "import sys; print(sys.implementation.cache_tag)")"; export PYTHONPATH="${AKA_GEAK_SDK_PATH}${PYTHONPATH:+:$PYTHONPATH}"; fi && if [[ "${AKA_ROCM_SDK_CORE_RUNTIME:-0}" == "1" ]]; then sdk_library_path="$(python src/scripts/rocm_sdk_runtime.py --print-library-path)" && export LD_LIBRARY_PATH="$sdk_library_path" && mkdir -p /tmp/aka-runtime-bin && ln -sf "$AGENT_KERNEL_ARENA_WORKDIR/src/scripts/rocm_sdk_runtime.py" /tmp/aka-runtime-bin/rocprofv3 && export PATH="/tmp/aka-runtime-bin:$PATH"; fi && if [[ "${AGENT_KERNEL_ARENA_ISOLATED_HOME:-0}" == "1" ]]; then bash src/scripts/docker_benchmark.sh _container_prepare_worker_home; fi && exec "$@"' _ "$@" } extract_config_name() { @@ -1591,7 +1612,7 @@ PY # pulls the SDK's full dependency closure (a few hundred MB, several minutes # on first run). This is a one-time provisioning cost — later runs import the # SDK via PYTHONPATH and short-circuit above. - local target="${PYTHONUSERBASE:-$PWD/.aka-pyuserbase}/geak-sdk" + local target="${AKA_GEAK_SDK_PATH:-${PYTHONUSERBASE:-$PWD/.aka-pyuserbase}/geak-sdk/$(python -c 'import sys; print(sys.implementation.cache_tag)')}" echo "claude-agent-sdk not found in image; installing into $target ..." python -m pip install --upgrade --target "$target" -r agents/geak/requirements.txt PYTHONPATH="$target${PYTHONPATH:+:$PYTHONPATH}" python -c 'import claude_agent_sdk; print("claude-agent-sdk=" + str(getattr(claude_agent_sdk, "__version__", "unknown")) + " setup OK")' diff --git a/src/scripts/rocm_sdk_runtime.py b/src/scripts/rocm_sdk_runtime.py new file mode 100755 index 000000000..50eb3f9f0 --- /dev/null +++ b/src/scripts/rocm_sdk_runtime.py @@ -0,0 +1,52 @@ +#!/usr/bin/env python3 +"""Keep ROCm SDK workloads and the profiler on the same core library tree.""" + +from __future__ import annotations + +import importlib.util +import os +from pathlib import Path +import sys + + +def sdk_core() -> Path: + spec = importlib.util.find_spec("_rocm_sdk_core") + if spec is None or spec.origin is None: + raise RuntimeError("The selected runtime has no ROCm SDK core package") + return Path(spec.origin).parent + + +def runtime_environment(environment: dict[str, str]) -> dict[str, str]: + core = sdk_core() + # The expanded developer tree can contain a second copy of COMGR/LLVM. + # PyTorch's absolute core preload cannot be redirected by LD_LIBRARY_PATH. + # Starting both the profiler and its workload with core libraries prevents + # loading that second copy and duplicate LLVM option registration. + core_lib = str(core / "lib") + devel_lib = str(core.parent / "_rocm_sdk_devel" / "lib") + libraries = [core_lib, *( + path for path in environment.get("LD_LIBRARY_PATH", "").split(":") + if path and path not in {core_lib, devel_lib} + )] + return { + **environment, "LD_LIBRARY_PATH": ":".join(libraries), + } + + +def profiler_command(arguments: list[str], environment: dict[str, str]): + executable = sdk_core() / "bin" / "rocprofv3" + if not executable.is_file(): + raise RuntimeError(f"ROCm SDK profiler is missing: {executable}") + return [str(executable), *arguments], runtime_environment(environment) + + +def main() -> None: + if sys.argv[1:] == ["--print-library-path"]: + print(runtime_environment(dict(os.environ))["LD_LIBRARY_PATH"]) + return + command, environment = profiler_command(sys.argv[1:], dict(os.environ)) + os.execve(command[0], command, environment) + + +if __name__ == "__main__": + main() diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/README.md b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/README.md index 618b85612..41a04c9bd 100644 --- a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/README.md +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/README.md @@ -67,3 +67,7 @@ also forbidden. Ordinary Python utilities, PyTorch allocation/layout operations, and the task's bundled `kernels/` helpers remain available under the existing numerical and timing contract. Baseline checks retain their declared initial backend; the final candidate must use FlyDSL. + +The task bundles the attributed legacy buffer/vector API adapters in +`flydsl_compat/` for FlyDSL 0.3 runtimes. Older runtimes use their installed +helpers. Kernel computation, inputs and numerical gates are unchanged. diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/LICENSE b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/LICENSE new file mode 100644 index 000000000..c73e2f2c4 --- /dev/null +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/LICENSE @@ -0,0 +1,17 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + Copyright 2025 FlyDSL Project Contributors + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/SOURCE.md b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/SOURCE.md new file mode 100644 index 000000000..83ea5ad77 --- /dev/null +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/SOURCE.md @@ -0,0 +1,13 @@ +# Pinned FlyDSL compatibility helpers + +Derived from ROCm/FlyDSL commit `28a18d328b4882c999864b2df2f8f9fe3fcc8b47`, +`python/flydsl/expr/{buffer_ops,vector,meta}.py`, under Apache-2.0 (see LICENSE). +The original buffer descriptor, byte offset, masking and cache policy is retained. +Relative dependency imports now target the installed package. Removed memref +pointer extraction uses the current typed iterator and LLVM pointer conversion. +Legacy runtimes continue using their installed original helpers. No runtime +monkeypatching, external repository imports or downloads are performed. + +Current ROCDL load/store operations receive the same cache-policy bits through +their `aux` attribute instead of a removed SSA operand. Vector unwrapping uses +the current signature while retaining the surrounding MLIR location context. diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/__init__.py b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/__init__.py new file mode 100644 index 000000000..1bbff9ffa --- /dev/null +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/__init__.py @@ -0,0 +1,13 @@ +"""Task-local legacy API fallback; never modifies the installed FlyDSL package.""" +try: + import flydsl.expr.buffer_ops as buffer_ops +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.buffer_ops": + raise + from . import buffer_ops +try: + import flydsl.expr.vector as vector +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.vector": + raise + from . import vector diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/buffer_ops.py b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/buffer_ops.py new file mode 100644 index 000000000..ab27dd821 --- /dev/null +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/buffer_ops.py @@ -0,0 +1,603 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""AMD Buffer Load/Store Operations - High-level Python API + +This module provides high-level Python wrappers for AMD CDNA3/CDNA4 buffer operations. +Buffer operations use a scalar base pointer and per-thread offsets for efficient memory access. + +Example: + >>> from flydsl._mlir_helpers import buffer_ops + >>> from flydsl._mlir_helpers import arith + >>> import _mlir.extras.types as T + >>> + >>> # Create buffer resource from memref + >>> rsrc = buffer_ops.create_buffer_resource(A) + >>> + >>> # Compute offset + >>> offset = row * arith.index(4096) + col + >>> + >>> # Buffer load (4xf32) + >>> data = buffer_ops.buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Buffer store + >>> buffer_ops.buffer_store(data, rsrc, offset) +""" + +from typing import Optional, Union + +from flydsl._mlir import ir +from flydsl._mlir.dialects import arith as std_arith +from flydsl._mlir.dialects import llvm, rocdl +from flydsl._mlir.extras import types as T +from flydsl.runtime.device import is_rdna_arch +from .meta import traced_op + + +def _get_buffer_flags(arch=None): + """Get AMD buffer resource descriptor (V#) flags word (bits 127:96). + + Constructs the 32-bit flags field for rocdl.make.buffer.rsrc, following the + same logic as LLVM's AMDGPUToROCDL makeBufferRsrc(): + https://github.com/llvm/llvm-project/blob/main/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp + + Bit layout (common to all architectures): + bits [11:0] - DST_SEL: ignored by raw buffer intrinsics + bits [14:12] - DATA_FORMAT: must be nonzero, 7 = float + bits [18:15] - NUM_FORMAT: must be nonzero, 4 = 32-bit + bit [19] - In nested heap (0) + bit [20] - Behavior on unmap (0 = return 0 / ignore) + bits [22:21] - Index stride for swizzles (0) + bit [23] - Add thread ID (0) + bit [24] - Reserved: must be 1 on RDNA, 0 on CDNA + bits [26:25] - Reserved (0) + bit [27] - Non-volatile (CDNA only, 0) + bits [29:28] - OOB_SELECT (RDNA only): 0=structured, 2=none, 3=check offset + bits [31:30] - Type (must be 0) + + CDNA (gfx9xx): (7 << 12) | (4 << 15) = 0x20070 + RDNA (gfx10+): (7 << 12) | (4 << 15) | (1 << 24) | (2 << 28) = 0x21020070 + - bit 24 set to 1 (required on RDNA) + - OOB_SELECT=2 (no bounds checking, matching LLVM boundsCheck=false) + """ + import os + + if arch is None: + arch = os.environ.get("FLYDSL_GPU_ARCH") + flags = (7 << 12) | (4 << 15) + if is_rdna_arch(arch): + flags |= 1 << 24 # reserved bit, must be 1 on RDNA + flags |= 2 << 28 # OOB_SELECT = 2 (no bounds checking) + return flags + + +__all__ = [ + "create_llvm_ptr", + "get_element_ptr", + "create_buffer_resource", + "create_buffer_resource_from_addr", + "buffer_load", + "buffer_store", + "BufferResourceDescriptor", + "extract_base_index", +] + + +def _unwrap_value(value): + """Recursively unwrap ArithValue or similar wrappers to get the actual MLIR value. + + Handles: + - FlyDSL ArithValue (has ._value) + - flyc DSL Numeric like fx.Int32 (has .ir_value() method) + - flyc ArithValue (is already ir.Value subclass) + """ + # DSL Numeric (Int32, Float32, etc.) — use ir_value() to materialize + if hasattr(value, "ir_value") and not isinstance(value, ir.Value): + return value.ir_value() + max_depth = 10 # Safety limit + depth = 0 + while depth < max_depth and not isinstance(value, ir.Value): + if hasattr(value, "_value"): + value = value._value + elif hasattr(value, "value"): + value = value.value + else: + break + depth += 1 + return value + + +def _create_i32_constant(value: int) -> ir.Value: + """Create i32 constant using standard MLIR arith dialect.""" + i32_type = T.i32() + if value > 0x7FFFFFFF: + value = int(value - 2**32) + attr = ir.IntegerAttr.get(i32_type, value) + op = std_arith.ConstantOp(i32_type, attr) + return _unwrap_value(op.result) + + +def _create_i16_constant(value: int) -> ir.Value: + """Create i16 constant using standard MLIR arith dialect.""" + i16_type = T.i16() + attr = ir.IntegerAttr.get(i16_type, value) + op = std_arith.ConstantOp(i16_type, attr) + return _unwrap_value(op.result) + + +def _create_i64_constant(value: int) -> ir.Value: + """Create i64 constant using standard MLIR arith dialect.""" + i64_type = T.i64() + attr = ir.IntegerAttr.get(i64_type, value) + op = std_arith.ConstantOp(i64_type, attr) + return _unwrap_value(op.result) + + +def create_llvm_ptr(value, address_space: int = 0) -> ir.Value: + """Create an LLVM pointer from an integer or index value.""" + value = _unwrap_value(value) + if isinstance(value.type, ir.IndexType): + i64_type = T.i64() + value = _unwrap_value(std_arith.IndexCastOp(i64_type, value).result) + ptr_type = ir.Type.parse(f"!llvm.ptr<{address_space}>") + return llvm.IntToPtrOp(ptr_type, value).result + + +def extract_base_index(tensor, address_space: int = 1) -> ir.Value: + """Extract the base address of a fly.memref as an index value. + + Inverse of :func:`create_llvm_ptr` (index -> ptr). Useful when ISA + requires a raw pointer instead of a buffer resource descriptor + (e.g. global_atomic_pk_add_bf16 on gfx942). + """ + from flydsl._mlir.dialects import fly as _fly + from flydsl._mlir.dialects import memref as _memref + + raw = _unwrap_value(tensor) + try: + ir.MemRefType(raw.type) + return _memref.extract_aligned_pointer_as_index(raw) + except ValueError: + pass + + # FlyDSL 0.3 uses a typed iterator instead of the removed extract op. + from flydsl.expr import get_iter, to_llvm_ptr + ptr = to_llvm_ptr(get_iter(raw)) + i64_val = llvm.PtrToIntOp(ir.IntegerType.get_signless(64), ptr).result + return _unwrap_value(std_arith.IndexCastOp(ir.IndexType.get(), i64_val).result) + + +def get_element_ptr( + base_ptr, + byte_offset: Union[int, ir.Value, None] = None, + static_byte_offset: int = 0, + elem_type: Optional[ir.Type] = None, + no_wrap_flags=None, +) -> ir.Value: + """Build an LLVM GEP from a base pointer plus byte offsets.""" + _gep_dynamic_index_sentinel = -(2**31) + + base_ptr = _unwrap_value(base_ptr) + if not isinstance(static_byte_offset, int): + raise TypeError(f"static_byte_offset must be int, got {type(static_byte_offset).__name__}") + if elem_type is None: + elem_type = T.i8() + elif callable(elem_type): + elem_type = elem_type() + + if byte_offset is None: + dynamic_indices = [] + raw_constant_indices = [int(static_byte_offset)] + elif isinstance(byte_offset, int): + dynamic_indices = [] + raw_constant_indices = [int(byte_offset) + int(static_byte_offset)] + else: + offset_val = _unwrap_value(byte_offset) + if isinstance(offset_val.type, ir.IndexType): + i64_type = T.i64() + offset_val = _unwrap_value(std_arith.IndexCastOp(i64_type, offset_val).result) + elif not isinstance(offset_val.type, ir.IntegerType): + raise TypeError("byte_offset must be int, index, or integer-typed MLIR value; " f"got {offset_val.type}") + + if static_byte_offset != 0: + static_type = offset_val.type + static_attr = ir.IntegerAttr.get(static_type, int(static_byte_offset)) + static_const = _unwrap_value(std_arith.ConstantOp(static_type, static_attr).result) + offset_val = _unwrap_value(std_arith.AddIOp(offset_val, static_const).result) + + dynamic_indices = [offset_val] + raw_constant_indices = [_gep_dynamic_index_sentinel] + + return llvm.GEPOp( + base_ptr.type, + base_ptr, + dynamic_indices, + raw_constant_indices, + elem_type, + no_wrap_flags, + ).result + + +class BufferResourceDescriptor: + """AMD Buffer Resource Descriptor + + A buffer resource descriptor contains: + - base_pointer: Scalar base pointer (wave-uniform, stored in SGPRs) + - stride: Stride for structured buffers (typically 0 for contiguous) + - num_records: Buffer size in bytes + - flags: Data format and access flags + + The descriptor is stored in a special LLVM pointer type (!llvm.ptr<8>) + """ + + def __init__(self, rsrc: ir.Value): + """Initialize with ROCDL resource descriptor value.""" + self.rsrc = rsrc + + @staticmethod + def from_memref( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + data_format: str = "f32", + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, + ) -> "BufferResourceDescriptor": + """Create buffer resource descriptor from memref. + + Args: + memref_val: Memref value to create descriptor for + stride: Stride in elements (0 for contiguous) + max_size: If True, use max buffer size for flexibility + num_records_bytes: Override buffer size (in BYTES) used by hardware OOB checking. + If provided, this takes precedence over `max_size`. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + data_format: Data format ('f32', 'f16', 'i32', etc.) + + Returns: + BufferResourceDescriptor instance + + Example: + >>> rsrc = BufferResourceDescriptor.from_memref(A) + """ + # Extract raw pointer from fly.memref. + raw_val = _unwrap_value(memref_val) + from flydsl._mlir.dialects import fly as _fly + + # Preserve byte offsets and descriptor policy with the current pointer API. + from flydsl.expr import get_iter, to_llvm_ptr + base_ptr = to_llvm_ptr(get_iter(raw_val)) + if base_byte_offset is not None: + base_ptr = get_element_ptr(base_ptr, byte_offset=base_byte_offset) + + # Create buffer resource descriptor + flags_val = _get_buffer_flags() + flags = _create_i32_constant(flags_val) + stride_val = _create_i16_constant(stride) + + def _num_records_from_memref_type() -> Optional[int]: + """Best-effort: derive logical buffer size (in bytes) from static memref type.""" + try: + mt = ir.MemRefType(_unwrap_value(memref_val).type) + shape = list(mt.shape) + if any(int(d) < 0 for d in shape): + return None + # Compute element size in bytes (scalar element type). + elem_t = mt.element_type + elem_bits = getattr(elem_t, "width", None) + if elem_bits is None: + return None + elem_bytes = int(elem_bits) // 8 + if elem_bytes <= 0: + return None + num_elems = 1 + for d in shape: + num_elems *= int(d) + return int(num_elems) * int(elem_bytes) + except Exception: + return None + + if num_records_bytes is not None: + # Caller-provided size in BYTES (preferred for exact hardware OOB behavior). + if isinstance(num_records_bytes, int): + nbytes = int(num_records_bytes) + if nbytes <= 0: + nbytes = 0 + # Descriptor uses i32 bytes; clamp to the max representable. + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(nbytes) + else: + v = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(v.type, ir.IntegerType) or v.type.width != 64: + if isinstance(v.type, ir.IndexType): + op = std_arith.IndexCastOp(i64_type, v) + else: + op = std_arith.ExtSIOp(i64_type, v) + v = _unwrap_value(op.result) + num_records = v + elif max_size: + # Use max for flexibility (hardware will check actual bounds) + # Note: FlyDSL's rocdl.make.buffer.rsrc requires i32, not i64 + num_records = _create_i64_constant(0xFFFFFFFF) # FALLBACK_MAX_SIZE + else: + # Use the logical memref size (in bytes) for hardware OOB checking. + nbytes = _num_records_from_memref_type() + if nbytes is None: + # Fall back to max-size if we can't infer statically. + num_records = _create_i64_constant(0xFFFFFFFF) + else: + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(int(nbytes)) + + # Create resource descriptor (returns !llvm.ptr<8>) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + rsrc = rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride_val, num_records, flags).result + + return BufferResourceDescriptor(rsrc) + + +def create_buffer_resource_from_addr( + addr_i64: ir.Value, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from a raw i64 device address. + + Useful when working with runtime pointer arrays (e.g. IPC-mapped addresses + or device-side pointer tables) where no fly.memref is available. + The full address is encoded as the buffer base; callers should pass + byte offset 0 to buffer_load / buffer_store. + + Args: + addr_i64: Raw 64-bit device address (i64 MLIR value). + num_records_bytes: Optional buffer size in bytes for hardware OOB checking. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>). + + Example: + >>> rsrc = create_buffer_resource_from_addr(raw_addr_i64) + >>> data = buffer_load(rsrc, i32_zero, vec_width=4, dtype=T.i32) + """ + addr_i64 = _unwrap_value(addr_i64) + ptr_type = ir.Type.parse("!llvm.ptr") + base_ptr = llvm.IntToPtrOp(ptr_type, addr_i64).result + flags = _create_i32_constant(_get_buffer_flags()) + stride = _create_i16_constant(0) + if num_records_bytes is None: + num_records = _create_i64_constant(0xFFFFFFFF) + elif isinstance(num_records_bytes, int): + nbytes = max(0, min(int(num_records_bytes), 0xFFFFFFFF)) + num_records = _create_i64_constant(nbytes) + else: + num_records = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(num_records.type, ir.IntegerType) or num_records.type.width != 64: + if isinstance(num_records.type, ir.IndexType): + num_records = _unwrap_value(std_arith.IndexCastOp(i64_type, num_records).result) + else: + num_records = _unwrap_value(std_arith.ExtSIOp(i64_type, num_records).result) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + return rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride, num_records, flags).result + + +@traced_op +def create_buffer_resource( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from memref. + + This is a simplified wrapper around BufferResourceDescriptor.from_memref() + that returns the raw ROCDL resource value. + + Args: + memref_val: Memref value + stride: Buffer stride (0 for contiguous) + max_size: Use maximum buffer size + num_records_bytes: Override buffer size in bytes. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>) + + Example: + >>> rsrc = create_buffer_resource(A) + >>> data = buffer_load(rsrc, offset) + """ + desc = BufferResourceDescriptor.from_memref( + memref_val, + stride, + max_size, + num_records_bytes=num_records_bytes, + base_byte_offset=base_byte_offset, + ) + return desc.rsrc + + +@traced_op +def buffer_load( + rsrc: ir.Value, + offset: ir.Value, + vec_width: int = 4, + dtype=None, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + soffset_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """AMD buffer load operation. + + Load data from global memory using buffer descriptor and offset. + Uses hardware-level bounds checking and vectorization. + + Args: + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + vec_width: Vector width (1, 2, or 4) + dtype: Element data type (None for f32, or ir.F32Type, etc.) + mask: Optional mask for predicated load (i1 type) + cache_modifier: Cache control flags (0 for default) + soffset_bytes: Optional scalar offset (in BYTES) added by the buffer instruction (soffset). + Use this to fold small constant deltas into the instruction instead of emitting + extra VGPR address arithmetic. + + Returns: + Loaded data (scalar or vector depending on vec_width) + + Example: + >>> # Load 4xf32 + >>> data = buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Load with mask + >>> data = buffer_load(rsrc, offset, vec_width=4, mask=valid) + """ + # Default dtype to f32 + if dtype is None: + dtype = T.f32() + # Accept DSL Numeric class (e.g. fx.Int32) as dtype: unwrap to ir.Type + elif hasattr(dtype, "ir_type"): + dtype = dtype.ir_type + + # Unwrap offset first (accept Python ints and DSL Numeric values). + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: Buffer load offset is in BYTES, not elements! + # For vec4xf32, each element is 4 bytes, so multiply offset by 4 + element_bytes = dtype.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create vector type + if vec_width == 1: + result_type = dtype + else: + result_type = ir.VectorType.get([vec_width], dtype) + + # Create instruction offset and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(soffset_bytes) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer load + load_op = rocdl.RawPtrBufferLoadOp( + result_type, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) + + return load_op.result + + +@traced_op +def buffer_store( + data: ir.Value, + rsrc: ir.Value, + offset: ir.Value, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + *, + soffset_bytes: Optional[Union[int, ir.Value]] = None, + offset_is_bytes: bool = False, +): + """AMD buffer store operation. + + Store data to global memory using buffer descriptor and offset. + + Args: + data: Data to store (scalar or vector) + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + mask: Optional mask for predicated store (i1 type) + cache_modifier: Cache control flags (0 for default) + + Example: + >>> buffer_store(data, rsrc, offset) + >>> + >>> # Store with mask + >>> buffer_store(data, rsrc, offset, mask=valid) + """ + # Unwrap all inputs (accept DSL Numeric values via ir_value()) + if hasattr(data, "ir_value"): + data = data.ir_value() + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + data = _unwrap_value(data) + rsrc = _unwrap_value(rsrc) + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: RawPtrBufferStoreOp offset is in BYTES. + # For backward compat, `buffer_store()` accepts element offsets by default + # and scales them to bytes. Set `offset_is_bytes=True` to skip scaling. + if not offset_is_bytes: + # Get element size from data type + data_type = data.type + if hasattr(data_type, "element_type"): # Vector type + element_type = data_type.element_type + else: # Scalar type + element_type = data_type + element_bytes = element_type.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create instruction offset (soffset) and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(int(soffset_bytes)) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer store + rocdl.RawPtrBufferStoreOp( + data, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/meta.py b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/meta.py new file mode 100644 index 000000000..23e4c849a --- /dev/null +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/meta.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +import inspect +from functools import wraps + +from flydsl._mlir import ir + + +def _to_raw_value(obj): + if isinstance(obj, ir.Value): + return obj + if isinstance(obj, type): + return obj + if hasattr(obj, "__extract_to_ir_values__"): + values = obj.__extract_to_ir_values__() + if len(values) != 1: + raise ValueError(f"Primitive function expects 1 value, got {len(values)}") + return values[0] + if isinstance(obj, tuple): + return tuple(_to_raw_value(e) for e in obj) + if isinstance(obj, list): + return [_to_raw_value(e) for e in obj] + return obj + + +def _flatten_args(args, kwargs): + new_args = tuple(_to_raw_value(a) for a in args) + new_kwargs = {k: _to_raw_value(v) if k not in ("loc", "ip") else v for k, v in kwargs.items()} + return new_args, new_kwargs + + +def _caller_location(depth=1): + """Build an MLIR Location from the Python call-site *depth* frames up.""" + frame = inspect.currentframe() + for _ in range(depth + 1): + if frame is not None: + frame = frame.f_back + if frame is None: + return ir.Location.unknown() + + info = inspect.getframeinfo(frame) + pos = getattr(info, "positions", None) + line = pos.lineno if pos is not None else info.lineno + col = (pos.col_offset or 0) if pos is not None else 0 + file_loc = ir.Location.file(info.filename, line, col) + + if info.code_context: + label = " ".join(ln.strip() for ln in info.code_context) + else: + label = info.function + return ir.Location.name(label, childLoc=file_loc) + + +def traced_op(op): + @wraps(op) + def wrapper(*args, **kwargs): + loc = kwargs.pop("loc", None) + if loc is None: + loc = _caller_location(depth=1) + args, kwargs = _flatten_args(args, kwargs) + with loc: + return op(*args, **kwargs) + + return wrapper diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/vector.py b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/vector.py new file mode 100644 index 000000000..29f8289c2 --- /dev/null +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/flydsl_compat/vector.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""Vector dialect helpers and re-exports. + +The ``Vector`` class itself lives in ``typing.py`` alongside other builtin +DSL types. This module re-exports it for convenience and provides thin +wrappers around upstream ``_mlir.dialects.vector`` ops. +""" + +from __future__ import annotations + +from flydsl._mlir import ir +from flydsl._mlir.dialects import vector as _vector + +# Re-export upstream dialect for ``from flydsl.expr import vector; vector.broadcast(...)`` +from flydsl._mlir.dialects.vector import * # noqa: F401,F403,E402 +from .meta import traced_op + +# Re-export Vector and friends so ``from flydsl.expr.vector import Vector`` works +from flydsl.expr.typing import ReductionOp, Vector, empty_like, full, full_like, ones_like, zeros_like # noqa: F401 + +# ═══════════════════════════════════════════════════════════════════════ +# Dialect helper wrappers (legacy, will be deprecated) +# Prefer using Vector methods or _mlir.dialects.vector directly. +# ═══════════════════════════════════════════════════════════════════════ + + +@traced_op +def from_elements(*args, loc=None, ip=None, **kwargs): + """Construct a vector from scalar elements, auto-unwrapping ArithValue wrappers.""" + from flydsl.expr import arith as _arith_ext + + if len(args) >= 2: + args = list(args) + elems = args[1] + if isinstance(elems, (list, tuple)): + args[1] = [_arith_ext.unwrap(v) for v in elems] + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + +@traced_op +def store(value, memref, indices, *, loc=None, ip=None, **kwargs): + """Vector store wrapper that accepts ArithValue/wrappers for value/indices.""" + from flydsl.expr import arith as _arith_ext + + return _vector.store( + _arith_ext.unwrap(value), + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + **kwargs, + ) + + +# ----------------------------------------------------------------------------- +# Thin wrappers for common op classes that otherwise require `.result` access. +# ----------------------------------------------------------------------------- + + +@traced_op +def extract(vector, static_position=None, dynamic_position=None, *, loc=None, ip=None): + """Wrapper around `vector.ExtractOp(...).result`. + + When only ``dynamic_position`` is supplied (without explicit + ``static_position``), each dynamic index needs a corresponding + ``kDynamic`` sentinel in the static attribute so the ODS builder + pairs them correctly. This wrapper fills in the sentinels + automatically. + """ + from flydsl.expr import arith as _arith_ext + + if static_position is None: + static_position = [] + if dynamic_position is None: + dynamic_position = [] + dynamic_position = [_arith_ext.unwrap(i, index=True) for i in dynamic_position] + + n_static = len(static_position) + n_dynamic = len(dynamic_position) + if n_dynamic > 0 and n_static < n_dynamic: + kDynamic = ir.ShapedType.get_dynamic_size() + static_position = list(static_position) + [kDynamic] * (n_dynamic - n_static) + + return _vector.ExtractOp( + _arith_ext.unwrap(vector), + static_position=static_position, + dynamic_position=dynamic_position, + loc=loc, + ip=ip, + ).result + + +@traced_op +def load_op(result_type, memref, indices, *, loc=None, ip=None): + """Wrapper around `vector.LoadOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.LoadOp( + result_type, + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + ).result + + +@traced_op +def bitcast(result_type, source, *, loc=None, ip=None): + """Wrapper around `vector.BitCastOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.BitCastOp( + result_type, + _arith_ext.unwrap(source), + loc=loc, + ip=ip, + ).result diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/kernel.py b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/kernel.py index c2517e37c..10e778288 100644 --- a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/kernel.py +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/kernel.py @@ -11,7 +11,8 @@ import flydsl.expr as fx from flydsl._mlir import ir from flydsl.compiler.kernel_function import CompilationContext -from flydsl.expr import arith, buffer_ops, const_expr, gpu, range_constexpr, rocdl, vector +from flydsl.expr import arith, const_expr, gpu, range_constexpr, rocdl +from flydsl_compat import buffer_ops, vector from flydsl.expr.typing import T from flydsl.expr.typing import Vector as Vec from flydsl.runtime.device import get_rocm_arch as get_hip_arch @@ -104,6 +105,11 @@ def _out_elem_type(): def _out_elem_dtype(): return fx.BFloat16 if is_bf16_out else fx.Float16 + def _fp8_elem_type(): + # Match the input encoding and the former architecture-dependent T.f8 + # alias. Construct the IR type only inside the compilation context. + return fx.Float8E4M3FN.ir_type if _is_gfx950 else fx.Float8E4M3FNUZ.ir_type + epilog_tag = "cshuffle" if use_cshuffle_epilog else "direct" module_name = (f"bs_gemm_{out_dtype}_{epilog_tag}" f"_t{tile_m}x{tile_n}x{tile_k}").replace("-", "_") @@ -175,8 +181,8 @@ def kernel_gemm( base_ptr_pong = allocator_pong.get_base() base_ptr_ping = allocator_ping.get_base() - lds_a_pong = SmemPtr(base_ptr_pong, lds_pong_offset, T.f8, shape=(tile_m * tile_k,)).get() - lds_a_ping = SmemPtr(base_ptr_ping, lds_ping_offset, T.f8, shape=(tile_m * tile_k,)).get() + lds_a_pong = SmemPtr(base_ptr_pong, lds_pong_offset, _fp8_elem_type(), shape=(tile_m * tile_k,)).get() + lds_a_ping = SmemPtr(base_ptr_ping, lds_ping_offset, _fp8_elem_type(), shape=(tile_m * tile_k,)).get() if const_expr(use_cshuffle_epilog): lds_out = SmemPtr(base_ptr_pong, lds_pong_offset, _out_elem_type(), shape=(tile_m * tile_n,)).get() @@ -239,7 +245,7 @@ def load_b_pack(base_k, ki_step, ni): n_blk=n_blk_list[ni], n_intra=n_intra_list[ni], lane_div_16=lane_div_16, - elem_type=T.f8, + elem_type=_fp8_elem_type(), kpack_bytes=kpack_bytes, elem_bytes=elem_bytes, ) @@ -259,7 +265,7 @@ def load_b_packs_k64(base_k, ku: int, ni: int): vector, b_rsrc, idx_pack, - elem_type=T.f8, + elem_type=_fp8_elem_type(), vec_elems=16, elem_bytes=elem_bytes, offset_in_bytes=True, @@ -285,7 +291,7 @@ def load_b_tile(base_k): def lds_load_16b(curr_row_a_lds, col_base, lds_buffer): col_base_swz = swizzle_xor16(curr_row_a_lds, col_base, k_blocks16) idx_a16 = curr_row_a_lds * _lds_k_dim_c + col_base_swz - return vector.load_op(T.f8x16, lds_buffer, [idx_a16]) + return vector.load_op(ir.VectorType.get([16], _fp8_elem_type()), lds_buffer, [idx_a16]) def lds_load_packs_k64(curr_row_a_lds, col_base, lds_buffer): loaded_a16 = lds_load_16b(curr_row_a_lds, col_base, lds_buffer) @@ -305,7 +311,7 @@ def load_a(idx_i32, a_load_bytes_v): return buffer_copy_gmem16_dwordx4( buffer_ops, vector, - elem_type=T.f8, + elem_type=_fp8_elem_type(), idx_i32=idx_i32, rsrc=a_rsrc, vec_elems=16, @@ -348,7 +354,7 @@ def store_a_tile_to_lds(vec_a_parts, lds_buffer, a_load_bytes_v, tx_i32_base_v, arith, vector, lds_memref=lds_buffer, - vec16_ty=T.f8x16, + vec16_ty=ir.VectorType.get([16], _fp8_elem_type()), layout_lds=layout_lds, row_local=row_a_local, col_local_i32=col_a_local_i32, @@ -363,7 +369,7 @@ def store_a_tile_to_lds(vec_a_parts, lds_buffer, a_load_bytes_v, tx_i32_base_v, arith, vector, lds_memref=lds_buffer, - vec8_ty=T.f8x8, + vec8_ty=ir.VectorType.get([8], _fp8_elem_type()), layout_lds=layout_lds, row_local=row_a_local, col_local_i32=col_a_local_i32, diff --git a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/kernels/kernels_common.py b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/kernels/kernels_common.py index 42058b6b1..411efb712 100644 --- a/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/kernels/kernels_common.py +++ b/tasks/flydsl2flydsl/blockscale_preshuffle_gemm_kernel/kernels/kernels_common.py @@ -16,7 +16,7 @@ from flydsl._mlir.dialects import gpu as _gpu from flydsl._mlir.dialects import llvm as _llvm from flydsl._mlir.dialects import scf as _scf -from flydsl.expr import buffer_ops +from flydsl_compat import buffer_ops from flydsl.expr.typing import T from flydsl.runtime.device import get_rocm_arch, is_rdna_arch diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/README.md b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/README.md index 09374600d..7dc4c72db 100644 --- a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/README.md +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/README.md @@ -6,8 +6,11 @@ always use this workspace's declared source; they never search other workspaces. Optimize the FlyDSL Paged Attention Decode FP8 kernel for AMD MI300X GPU. The kernel implements paged KV-cache attention decode with FP8 quantized -keys/values, MFMA-based dot products, online softmax, multi-partition -reduce, and supports both one-shot and split-reduce modes. +keys/values, MFMA-based dot products, online softmax, and split-partial +reduction. The eight declared workloads evaluate persistent scheduling with a +1027-token context and no sliding window. Other launch routes present in the +source are outside this task's measured workload; these cases do not qualify a +one-shot or sliding-window path. You MUST keep the kernel in FlyDSL — do NOT rewrite it in HIP, CUDA, or Triton. Only `candidate.editable` files in `config.yaml` may be changed. Preserve the declared @@ -88,3 +91,7 @@ helpers may only be called directly in those functions in `kernel.py`, not exported, introspected, or used as access to another AITER operator. Other AITER operators and namespace imports remain forbidden. This preserves the existing metadata/reduce glue and does not permit delegating attention computation. + +The task bundles the attributed legacy buffer/vector API adapters in +`flydsl_compat/` for FlyDSL 0.3 runtimes. Older runtimes use their installed +helpers. Kernel computation, inputs and numerical gates are unchanged. diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/config.yaml b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/config.yaml index 5a843e513..3660772da 100644 --- a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/config.yaml +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/config.yaml @@ -1,7 +1,7 @@ schema_version: 2 description: "Optimize the FlyDSL Paged Attention Decode FP8 kernel for AMD MI300X GPU.\nThe kernel implements\ \ paged KV-cache attention decode with FP8 quantized\nkeys/values, MFMA-based dot products, online softmax,\ - \ multi-partition\nreduce, and supports both one-shot and split-reduce modes.\nYou MUST keep the kernel\ + \ multi-partition\nreduce. The declared workload uses persistent scheduling with split-partial reduction, a 1027-token context, and no sliding window.\nYou MUST keep the kernel\ \ in FlyDSL \u2014 do NOT rewrite it in HIP, CUDA, or Triton." instructions: - README.md diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/LICENSE b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/LICENSE new file mode 100644 index 000000000..c73e2f2c4 --- /dev/null +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/LICENSE @@ -0,0 +1,17 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + Copyright 2025 FlyDSL Project Contributors + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/SOURCE.md b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/SOURCE.md new file mode 100644 index 000000000..83ea5ad77 --- /dev/null +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/SOURCE.md @@ -0,0 +1,13 @@ +# Pinned FlyDSL compatibility helpers + +Derived from ROCm/FlyDSL commit `28a18d328b4882c999864b2df2f8f9fe3fcc8b47`, +`python/flydsl/expr/{buffer_ops,vector,meta}.py`, under Apache-2.0 (see LICENSE). +The original buffer descriptor, byte offset, masking and cache policy is retained. +Relative dependency imports now target the installed package. Removed memref +pointer extraction uses the current typed iterator and LLVM pointer conversion. +Legacy runtimes continue using their installed original helpers. No runtime +monkeypatching, external repository imports or downloads are performed. + +Current ROCDL load/store operations receive the same cache-policy bits through +their `aux` attribute instead of a removed SSA operand. Vector unwrapping uses +the current signature while retaining the surrounding MLIR location context. diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/__init__.py b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/__init__.py new file mode 100644 index 000000000..1bbff9ffa --- /dev/null +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/__init__.py @@ -0,0 +1,13 @@ +"""Task-local legacy API fallback; never modifies the installed FlyDSL package.""" +try: + import flydsl.expr.buffer_ops as buffer_ops +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.buffer_ops": + raise + from . import buffer_ops +try: + import flydsl.expr.vector as vector +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.vector": + raise + from . import vector diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/buffer_ops.py b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/buffer_ops.py new file mode 100644 index 000000000..ab27dd821 --- /dev/null +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/buffer_ops.py @@ -0,0 +1,603 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""AMD Buffer Load/Store Operations - High-level Python API + +This module provides high-level Python wrappers for AMD CDNA3/CDNA4 buffer operations. +Buffer operations use a scalar base pointer and per-thread offsets for efficient memory access. + +Example: + >>> from flydsl._mlir_helpers import buffer_ops + >>> from flydsl._mlir_helpers import arith + >>> import _mlir.extras.types as T + >>> + >>> # Create buffer resource from memref + >>> rsrc = buffer_ops.create_buffer_resource(A) + >>> + >>> # Compute offset + >>> offset = row * arith.index(4096) + col + >>> + >>> # Buffer load (4xf32) + >>> data = buffer_ops.buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Buffer store + >>> buffer_ops.buffer_store(data, rsrc, offset) +""" + +from typing import Optional, Union + +from flydsl._mlir import ir +from flydsl._mlir.dialects import arith as std_arith +from flydsl._mlir.dialects import llvm, rocdl +from flydsl._mlir.extras import types as T +from flydsl.runtime.device import is_rdna_arch +from .meta import traced_op + + +def _get_buffer_flags(arch=None): + """Get AMD buffer resource descriptor (V#) flags word (bits 127:96). + + Constructs the 32-bit flags field for rocdl.make.buffer.rsrc, following the + same logic as LLVM's AMDGPUToROCDL makeBufferRsrc(): + https://github.com/llvm/llvm-project/blob/main/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp + + Bit layout (common to all architectures): + bits [11:0] - DST_SEL: ignored by raw buffer intrinsics + bits [14:12] - DATA_FORMAT: must be nonzero, 7 = float + bits [18:15] - NUM_FORMAT: must be nonzero, 4 = 32-bit + bit [19] - In nested heap (0) + bit [20] - Behavior on unmap (0 = return 0 / ignore) + bits [22:21] - Index stride for swizzles (0) + bit [23] - Add thread ID (0) + bit [24] - Reserved: must be 1 on RDNA, 0 on CDNA + bits [26:25] - Reserved (0) + bit [27] - Non-volatile (CDNA only, 0) + bits [29:28] - OOB_SELECT (RDNA only): 0=structured, 2=none, 3=check offset + bits [31:30] - Type (must be 0) + + CDNA (gfx9xx): (7 << 12) | (4 << 15) = 0x20070 + RDNA (gfx10+): (7 << 12) | (4 << 15) | (1 << 24) | (2 << 28) = 0x21020070 + - bit 24 set to 1 (required on RDNA) + - OOB_SELECT=2 (no bounds checking, matching LLVM boundsCheck=false) + """ + import os + + if arch is None: + arch = os.environ.get("FLYDSL_GPU_ARCH") + flags = (7 << 12) | (4 << 15) + if is_rdna_arch(arch): + flags |= 1 << 24 # reserved bit, must be 1 on RDNA + flags |= 2 << 28 # OOB_SELECT = 2 (no bounds checking) + return flags + + +__all__ = [ + "create_llvm_ptr", + "get_element_ptr", + "create_buffer_resource", + "create_buffer_resource_from_addr", + "buffer_load", + "buffer_store", + "BufferResourceDescriptor", + "extract_base_index", +] + + +def _unwrap_value(value): + """Recursively unwrap ArithValue or similar wrappers to get the actual MLIR value. + + Handles: + - FlyDSL ArithValue (has ._value) + - flyc DSL Numeric like fx.Int32 (has .ir_value() method) + - flyc ArithValue (is already ir.Value subclass) + """ + # DSL Numeric (Int32, Float32, etc.) — use ir_value() to materialize + if hasattr(value, "ir_value") and not isinstance(value, ir.Value): + return value.ir_value() + max_depth = 10 # Safety limit + depth = 0 + while depth < max_depth and not isinstance(value, ir.Value): + if hasattr(value, "_value"): + value = value._value + elif hasattr(value, "value"): + value = value.value + else: + break + depth += 1 + return value + + +def _create_i32_constant(value: int) -> ir.Value: + """Create i32 constant using standard MLIR arith dialect.""" + i32_type = T.i32() + if value > 0x7FFFFFFF: + value = int(value - 2**32) + attr = ir.IntegerAttr.get(i32_type, value) + op = std_arith.ConstantOp(i32_type, attr) + return _unwrap_value(op.result) + + +def _create_i16_constant(value: int) -> ir.Value: + """Create i16 constant using standard MLIR arith dialect.""" + i16_type = T.i16() + attr = ir.IntegerAttr.get(i16_type, value) + op = std_arith.ConstantOp(i16_type, attr) + return _unwrap_value(op.result) + + +def _create_i64_constant(value: int) -> ir.Value: + """Create i64 constant using standard MLIR arith dialect.""" + i64_type = T.i64() + attr = ir.IntegerAttr.get(i64_type, value) + op = std_arith.ConstantOp(i64_type, attr) + return _unwrap_value(op.result) + + +def create_llvm_ptr(value, address_space: int = 0) -> ir.Value: + """Create an LLVM pointer from an integer or index value.""" + value = _unwrap_value(value) + if isinstance(value.type, ir.IndexType): + i64_type = T.i64() + value = _unwrap_value(std_arith.IndexCastOp(i64_type, value).result) + ptr_type = ir.Type.parse(f"!llvm.ptr<{address_space}>") + return llvm.IntToPtrOp(ptr_type, value).result + + +def extract_base_index(tensor, address_space: int = 1) -> ir.Value: + """Extract the base address of a fly.memref as an index value. + + Inverse of :func:`create_llvm_ptr` (index -> ptr). Useful when ISA + requires a raw pointer instead of a buffer resource descriptor + (e.g. global_atomic_pk_add_bf16 on gfx942). + """ + from flydsl._mlir.dialects import fly as _fly + from flydsl._mlir.dialects import memref as _memref + + raw = _unwrap_value(tensor) + try: + ir.MemRefType(raw.type) + return _memref.extract_aligned_pointer_as_index(raw) + except ValueError: + pass + + # FlyDSL 0.3 uses a typed iterator instead of the removed extract op. + from flydsl.expr import get_iter, to_llvm_ptr + ptr = to_llvm_ptr(get_iter(raw)) + i64_val = llvm.PtrToIntOp(ir.IntegerType.get_signless(64), ptr).result + return _unwrap_value(std_arith.IndexCastOp(ir.IndexType.get(), i64_val).result) + + +def get_element_ptr( + base_ptr, + byte_offset: Union[int, ir.Value, None] = None, + static_byte_offset: int = 0, + elem_type: Optional[ir.Type] = None, + no_wrap_flags=None, +) -> ir.Value: + """Build an LLVM GEP from a base pointer plus byte offsets.""" + _gep_dynamic_index_sentinel = -(2**31) + + base_ptr = _unwrap_value(base_ptr) + if not isinstance(static_byte_offset, int): + raise TypeError(f"static_byte_offset must be int, got {type(static_byte_offset).__name__}") + if elem_type is None: + elem_type = T.i8() + elif callable(elem_type): + elem_type = elem_type() + + if byte_offset is None: + dynamic_indices = [] + raw_constant_indices = [int(static_byte_offset)] + elif isinstance(byte_offset, int): + dynamic_indices = [] + raw_constant_indices = [int(byte_offset) + int(static_byte_offset)] + else: + offset_val = _unwrap_value(byte_offset) + if isinstance(offset_val.type, ir.IndexType): + i64_type = T.i64() + offset_val = _unwrap_value(std_arith.IndexCastOp(i64_type, offset_val).result) + elif not isinstance(offset_val.type, ir.IntegerType): + raise TypeError("byte_offset must be int, index, or integer-typed MLIR value; " f"got {offset_val.type}") + + if static_byte_offset != 0: + static_type = offset_val.type + static_attr = ir.IntegerAttr.get(static_type, int(static_byte_offset)) + static_const = _unwrap_value(std_arith.ConstantOp(static_type, static_attr).result) + offset_val = _unwrap_value(std_arith.AddIOp(offset_val, static_const).result) + + dynamic_indices = [offset_val] + raw_constant_indices = [_gep_dynamic_index_sentinel] + + return llvm.GEPOp( + base_ptr.type, + base_ptr, + dynamic_indices, + raw_constant_indices, + elem_type, + no_wrap_flags, + ).result + + +class BufferResourceDescriptor: + """AMD Buffer Resource Descriptor + + A buffer resource descriptor contains: + - base_pointer: Scalar base pointer (wave-uniform, stored in SGPRs) + - stride: Stride for structured buffers (typically 0 for contiguous) + - num_records: Buffer size in bytes + - flags: Data format and access flags + + The descriptor is stored in a special LLVM pointer type (!llvm.ptr<8>) + """ + + def __init__(self, rsrc: ir.Value): + """Initialize with ROCDL resource descriptor value.""" + self.rsrc = rsrc + + @staticmethod + def from_memref( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + data_format: str = "f32", + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, + ) -> "BufferResourceDescriptor": + """Create buffer resource descriptor from memref. + + Args: + memref_val: Memref value to create descriptor for + stride: Stride in elements (0 for contiguous) + max_size: If True, use max buffer size for flexibility + num_records_bytes: Override buffer size (in BYTES) used by hardware OOB checking. + If provided, this takes precedence over `max_size`. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + data_format: Data format ('f32', 'f16', 'i32', etc.) + + Returns: + BufferResourceDescriptor instance + + Example: + >>> rsrc = BufferResourceDescriptor.from_memref(A) + """ + # Extract raw pointer from fly.memref. + raw_val = _unwrap_value(memref_val) + from flydsl._mlir.dialects import fly as _fly + + # Preserve byte offsets and descriptor policy with the current pointer API. + from flydsl.expr import get_iter, to_llvm_ptr + base_ptr = to_llvm_ptr(get_iter(raw_val)) + if base_byte_offset is not None: + base_ptr = get_element_ptr(base_ptr, byte_offset=base_byte_offset) + + # Create buffer resource descriptor + flags_val = _get_buffer_flags() + flags = _create_i32_constant(flags_val) + stride_val = _create_i16_constant(stride) + + def _num_records_from_memref_type() -> Optional[int]: + """Best-effort: derive logical buffer size (in bytes) from static memref type.""" + try: + mt = ir.MemRefType(_unwrap_value(memref_val).type) + shape = list(mt.shape) + if any(int(d) < 0 for d in shape): + return None + # Compute element size in bytes (scalar element type). + elem_t = mt.element_type + elem_bits = getattr(elem_t, "width", None) + if elem_bits is None: + return None + elem_bytes = int(elem_bits) // 8 + if elem_bytes <= 0: + return None + num_elems = 1 + for d in shape: + num_elems *= int(d) + return int(num_elems) * int(elem_bytes) + except Exception: + return None + + if num_records_bytes is not None: + # Caller-provided size in BYTES (preferred for exact hardware OOB behavior). + if isinstance(num_records_bytes, int): + nbytes = int(num_records_bytes) + if nbytes <= 0: + nbytes = 0 + # Descriptor uses i32 bytes; clamp to the max representable. + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(nbytes) + else: + v = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(v.type, ir.IntegerType) or v.type.width != 64: + if isinstance(v.type, ir.IndexType): + op = std_arith.IndexCastOp(i64_type, v) + else: + op = std_arith.ExtSIOp(i64_type, v) + v = _unwrap_value(op.result) + num_records = v + elif max_size: + # Use max for flexibility (hardware will check actual bounds) + # Note: FlyDSL's rocdl.make.buffer.rsrc requires i32, not i64 + num_records = _create_i64_constant(0xFFFFFFFF) # FALLBACK_MAX_SIZE + else: + # Use the logical memref size (in bytes) for hardware OOB checking. + nbytes = _num_records_from_memref_type() + if nbytes is None: + # Fall back to max-size if we can't infer statically. + num_records = _create_i64_constant(0xFFFFFFFF) + else: + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(int(nbytes)) + + # Create resource descriptor (returns !llvm.ptr<8>) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + rsrc = rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride_val, num_records, flags).result + + return BufferResourceDescriptor(rsrc) + + +def create_buffer_resource_from_addr( + addr_i64: ir.Value, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from a raw i64 device address. + + Useful when working with runtime pointer arrays (e.g. IPC-mapped addresses + or device-side pointer tables) where no fly.memref is available. + The full address is encoded as the buffer base; callers should pass + byte offset 0 to buffer_load / buffer_store. + + Args: + addr_i64: Raw 64-bit device address (i64 MLIR value). + num_records_bytes: Optional buffer size in bytes for hardware OOB checking. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>). + + Example: + >>> rsrc = create_buffer_resource_from_addr(raw_addr_i64) + >>> data = buffer_load(rsrc, i32_zero, vec_width=4, dtype=T.i32) + """ + addr_i64 = _unwrap_value(addr_i64) + ptr_type = ir.Type.parse("!llvm.ptr") + base_ptr = llvm.IntToPtrOp(ptr_type, addr_i64).result + flags = _create_i32_constant(_get_buffer_flags()) + stride = _create_i16_constant(0) + if num_records_bytes is None: + num_records = _create_i64_constant(0xFFFFFFFF) + elif isinstance(num_records_bytes, int): + nbytes = max(0, min(int(num_records_bytes), 0xFFFFFFFF)) + num_records = _create_i64_constant(nbytes) + else: + num_records = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(num_records.type, ir.IntegerType) or num_records.type.width != 64: + if isinstance(num_records.type, ir.IndexType): + num_records = _unwrap_value(std_arith.IndexCastOp(i64_type, num_records).result) + else: + num_records = _unwrap_value(std_arith.ExtSIOp(i64_type, num_records).result) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + return rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride, num_records, flags).result + + +@traced_op +def create_buffer_resource( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from memref. + + This is a simplified wrapper around BufferResourceDescriptor.from_memref() + that returns the raw ROCDL resource value. + + Args: + memref_val: Memref value + stride: Buffer stride (0 for contiguous) + max_size: Use maximum buffer size + num_records_bytes: Override buffer size in bytes. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>) + + Example: + >>> rsrc = create_buffer_resource(A) + >>> data = buffer_load(rsrc, offset) + """ + desc = BufferResourceDescriptor.from_memref( + memref_val, + stride, + max_size, + num_records_bytes=num_records_bytes, + base_byte_offset=base_byte_offset, + ) + return desc.rsrc + + +@traced_op +def buffer_load( + rsrc: ir.Value, + offset: ir.Value, + vec_width: int = 4, + dtype=None, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + soffset_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """AMD buffer load operation. + + Load data from global memory using buffer descriptor and offset. + Uses hardware-level bounds checking and vectorization. + + Args: + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + vec_width: Vector width (1, 2, or 4) + dtype: Element data type (None for f32, or ir.F32Type, etc.) + mask: Optional mask for predicated load (i1 type) + cache_modifier: Cache control flags (0 for default) + soffset_bytes: Optional scalar offset (in BYTES) added by the buffer instruction (soffset). + Use this to fold small constant deltas into the instruction instead of emitting + extra VGPR address arithmetic. + + Returns: + Loaded data (scalar or vector depending on vec_width) + + Example: + >>> # Load 4xf32 + >>> data = buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Load with mask + >>> data = buffer_load(rsrc, offset, vec_width=4, mask=valid) + """ + # Default dtype to f32 + if dtype is None: + dtype = T.f32() + # Accept DSL Numeric class (e.g. fx.Int32) as dtype: unwrap to ir.Type + elif hasattr(dtype, "ir_type"): + dtype = dtype.ir_type + + # Unwrap offset first (accept Python ints and DSL Numeric values). + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: Buffer load offset is in BYTES, not elements! + # For vec4xf32, each element is 4 bytes, so multiply offset by 4 + element_bytes = dtype.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create vector type + if vec_width == 1: + result_type = dtype + else: + result_type = ir.VectorType.get([vec_width], dtype) + + # Create instruction offset and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(soffset_bytes) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer load + load_op = rocdl.RawPtrBufferLoadOp( + result_type, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) + + return load_op.result + + +@traced_op +def buffer_store( + data: ir.Value, + rsrc: ir.Value, + offset: ir.Value, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + *, + soffset_bytes: Optional[Union[int, ir.Value]] = None, + offset_is_bytes: bool = False, +): + """AMD buffer store operation. + + Store data to global memory using buffer descriptor and offset. + + Args: + data: Data to store (scalar or vector) + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + mask: Optional mask for predicated store (i1 type) + cache_modifier: Cache control flags (0 for default) + + Example: + >>> buffer_store(data, rsrc, offset) + >>> + >>> # Store with mask + >>> buffer_store(data, rsrc, offset, mask=valid) + """ + # Unwrap all inputs (accept DSL Numeric values via ir_value()) + if hasattr(data, "ir_value"): + data = data.ir_value() + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + data = _unwrap_value(data) + rsrc = _unwrap_value(rsrc) + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: RawPtrBufferStoreOp offset is in BYTES. + # For backward compat, `buffer_store()` accepts element offsets by default + # and scales them to bytes. Set `offset_is_bytes=True` to skip scaling. + if not offset_is_bytes: + # Get element size from data type + data_type = data.type + if hasattr(data_type, "element_type"): # Vector type + element_type = data_type.element_type + else: # Scalar type + element_type = data_type + element_bytes = element_type.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create instruction offset (soffset) and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(int(soffset_bytes)) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer store + rocdl.RawPtrBufferStoreOp( + data, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/meta.py b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/meta.py new file mode 100644 index 000000000..23e4c849a --- /dev/null +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/meta.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +import inspect +from functools import wraps + +from flydsl._mlir import ir + + +def _to_raw_value(obj): + if isinstance(obj, ir.Value): + return obj + if isinstance(obj, type): + return obj + if hasattr(obj, "__extract_to_ir_values__"): + values = obj.__extract_to_ir_values__() + if len(values) != 1: + raise ValueError(f"Primitive function expects 1 value, got {len(values)}") + return values[0] + if isinstance(obj, tuple): + return tuple(_to_raw_value(e) for e in obj) + if isinstance(obj, list): + return [_to_raw_value(e) for e in obj] + return obj + + +def _flatten_args(args, kwargs): + new_args = tuple(_to_raw_value(a) for a in args) + new_kwargs = {k: _to_raw_value(v) if k not in ("loc", "ip") else v for k, v in kwargs.items()} + return new_args, new_kwargs + + +def _caller_location(depth=1): + """Build an MLIR Location from the Python call-site *depth* frames up.""" + frame = inspect.currentframe() + for _ in range(depth + 1): + if frame is not None: + frame = frame.f_back + if frame is None: + return ir.Location.unknown() + + info = inspect.getframeinfo(frame) + pos = getattr(info, "positions", None) + line = pos.lineno if pos is not None else info.lineno + col = (pos.col_offset or 0) if pos is not None else 0 + file_loc = ir.Location.file(info.filename, line, col) + + if info.code_context: + label = " ".join(ln.strip() for ln in info.code_context) + else: + label = info.function + return ir.Location.name(label, childLoc=file_loc) + + +def traced_op(op): + @wraps(op) + def wrapper(*args, **kwargs): + loc = kwargs.pop("loc", None) + if loc is None: + loc = _caller_location(depth=1) + args, kwargs = _flatten_args(args, kwargs) + with loc: + return op(*args, **kwargs) + + return wrapper diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/vector.py b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/vector.py new file mode 100644 index 000000000..29f8289c2 --- /dev/null +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/flydsl_compat/vector.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""Vector dialect helpers and re-exports. + +The ``Vector`` class itself lives in ``typing.py`` alongside other builtin +DSL types. This module re-exports it for convenience and provides thin +wrappers around upstream ``_mlir.dialects.vector`` ops. +""" + +from __future__ import annotations + +from flydsl._mlir import ir +from flydsl._mlir.dialects import vector as _vector + +# Re-export upstream dialect for ``from flydsl.expr import vector; vector.broadcast(...)`` +from flydsl._mlir.dialects.vector import * # noqa: F401,F403,E402 +from .meta import traced_op + +# Re-export Vector and friends so ``from flydsl.expr.vector import Vector`` works +from flydsl.expr.typing import ReductionOp, Vector, empty_like, full, full_like, ones_like, zeros_like # noqa: F401 + +# ═══════════════════════════════════════════════════════════════════════ +# Dialect helper wrappers (legacy, will be deprecated) +# Prefer using Vector methods or _mlir.dialects.vector directly. +# ═══════════════════════════════════════════════════════════════════════ + + +@traced_op +def from_elements(*args, loc=None, ip=None, **kwargs): + """Construct a vector from scalar elements, auto-unwrapping ArithValue wrappers.""" + from flydsl.expr import arith as _arith_ext + + if len(args) >= 2: + args = list(args) + elems = args[1] + if isinstance(elems, (list, tuple)): + args[1] = [_arith_ext.unwrap(v) for v in elems] + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + +@traced_op +def store(value, memref, indices, *, loc=None, ip=None, **kwargs): + """Vector store wrapper that accepts ArithValue/wrappers for value/indices.""" + from flydsl.expr import arith as _arith_ext + + return _vector.store( + _arith_ext.unwrap(value), + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + **kwargs, + ) + + +# ----------------------------------------------------------------------------- +# Thin wrappers for common op classes that otherwise require `.result` access. +# ----------------------------------------------------------------------------- + + +@traced_op +def extract(vector, static_position=None, dynamic_position=None, *, loc=None, ip=None): + """Wrapper around `vector.ExtractOp(...).result`. + + When only ``dynamic_position`` is supplied (without explicit + ``static_position``), each dynamic index needs a corresponding + ``kDynamic`` sentinel in the static attribute so the ODS builder + pairs them correctly. This wrapper fills in the sentinels + automatically. + """ + from flydsl.expr import arith as _arith_ext + + if static_position is None: + static_position = [] + if dynamic_position is None: + dynamic_position = [] + dynamic_position = [_arith_ext.unwrap(i, index=True) for i in dynamic_position] + + n_static = len(static_position) + n_dynamic = len(dynamic_position) + if n_dynamic > 0 and n_static < n_dynamic: + kDynamic = ir.ShapedType.get_dynamic_size() + static_position = list(static_position) + [kDynamic] * (n_dynamic - n_static) + + return _vector.ExtractOp( + _arith_ext.unwrap(vector), + static_position=static_position, + dynamic_position=dynamic_position, + loc=loc, + ip=ip, + ).result + + +@traced_op +def load_op(result_type, memref, indices, *, loc=None, ip=None): + """Wrapper around `vector.LoadOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.LoadOp( + result_type, + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + ).result + + +@traced_op +def bitcast(result_type, source, *, loc=None, ip=None): + """Wrapper around `vector.BitCastOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.BitCastOp( + result_type, + _arith_ext.unwrap(source), + loc=loc, + ip=ip, + ).result diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/kernel.py b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/kernel.py index 90c05b3df..3648a2366 100644 --- a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/kernel.py +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/kernel.py @@ -24,7 +24,8 @@ from flydsl._mlir import ir from flydsl._mlir.dialects import llvm from flydsl.compiler.kernel_function import CompilationContext -from flydsl.expr import arith, buffer_ops, const_expr, gpu, range_constexpr, rocdl, vector +from flydsl.expr import arith, const_expr, gpu, range_constexpr, rocdl +from flydsl_compat import buffer_ops, vector from flydsl.expr import math as fly_math from flydsl.expr.typing import Int32, T from flydsl.runtime.device import get_rocm_arch as get_hip_arch diff --git a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/kernels/pa_decode_swa.py b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/kernels/pa_decode_swa.py index cbeff4eb4..730c469f6 100644 --- a/tasks/flydsl2flydsl/pa_decode_fp8_kernel/kernels/pa_decode_swa.py +++ b/tasks/flydsl2flydsl/pa_decode_fp8_kernel/kernels/pa_decode_swa.py @@ -12,7 +12,8 @@ from flydsl._mlir import ir from flydsl._mlir.dialects import llvm from flydsl.compiler.kernel_function import CompilationContext -from flydsl.expr import arith, buffer_ops, const_expr, gpu, range_constexpr, rocdl, vector +from flydsl.expr import arith, const_expr, gpu, range_constexpr, rocdl +from flydsl_compat import buffer_ops, vector from flydsl.expr import math as fly_math from flydsl.expr.typing import Int32, T from flydsl.runtime.device import get_rocm_arch as get_hip_arch diff --git a/tasks/torch2flydsl/batched_gemm_a8w8_kernel/README.md b/tasks/torch2flydsl/batched_gemm_a8w8_kernel/README.md index 9aae166e6..d46e0b931 100644 --- a/tasks/torch2flydsl/batched_gemm_a8w8_kernel/README.md +++ b/tasks/torch2flydsl/batched_gemm_a8w8_kernel/README.md @@ -78,3 +78,14 @@ outputs through the canonical collector, with checks outside timed windows. Candidate-only dispatch auditing permits preparation allocations/copies/views, requires a FlyDSL runtime launch, and rejects AITER/operator compute shortcuts; reference and baseline calls cannot satisfy the candidate launch requirement. + +The action adapter reports the measured Event settings under +`metadata.device_timing` and derives `metadata.timed_output_checked` from the +harness's measured-output check. This lets the validator establish that graph +replay is not applicable without changing numerical or timing checks. + +Candidate code must not inspect or mutate the evaluator's loaded Python modules, +function internals, frames, or tracing hooks. The dependency check rejects module +registry access and dynamic introspection before candidate import, including +aliased imports. This protects the comparator from the known module-state bypass; +it does not make same-process Python execution a security sandbox. diff --git a/tasks/torch2flydsl/batched_gemm_a8w8_kernel/scripts/task_actions.py b/tasks/torch2flydsl/batched_gemm_a8w8_kernel/scripts/task_actions.py index 520decc8f..d1beb9e98 100644 --- a/tasks/torch2flydsl/batched_gemm_a8w8_kernel/scripts/task_actions.py +++ b/tasks/torch2flydsl/batched_gemm_a8w8_kernel/scripts/task_actions.py @@ -15,4 +15,12 @@ def check(h): return sorted(observed) def performance(h): - return h.arena_benchmark(warmup=10, iters=100, verbose=True) + rows = h.arena_benchmark(warmup=10, iters=100, verbose=True) + for row in rows: + # Preserve the measured values and expose the existing post-timing + # checks through the validator's event-timing evidence contract. + row["device_timing"] = { + key: value for key, value in row.items() if key.startswith("benchmark_") + } + row["timed_output_checked"] = row.get("timed_output_correctness") == "PASS" + return rows diff --git a/tasks/torch2flydsl/batched_gemm_a8w8_kernel/task_runtime.py b/tasks/torch2flydsl/batched_gemm_a8w8_kernel/task_runtime.py index 72052e727..4ce57890a 100644 --- a/tasks/torch2flydsl/batched_gemm_a8w8_kernel/task_runtime.py +++ b/tasks/torch2flydsl/batched_gemm_a8w8_kernel/task_runtime.py @@ -86,7 +86,7 @@ def source_state(cfg): def check_dependencies(paths, final_language=True): """Enforce declared implementation dependencies, including from X import Y.""" - forbidden = {"src", "agents", "model", "test_kernel_harness", "task_runtime", "task_reference", "task_baseline", "reference_controls", "scripts"} + forbidden = {"src", "agents", "model", "test_kernel_harness", "task_runtime", "task_reference", "task_baseline", "reference_controls", "scripts", "inspect", "gc", "builtins", "importlib"} backend_seen = False for path in paths: tree = ast.parse(path.read_text(), filename=str(path)) @@ -106,6 +106,8 @@ def check_dependencies(paths, final_language=True): parts = set(module.split(".")) if parts & forbidden: raise ValueError(f"Protected dependency in candidate: {module}") + if module in {"sys.modules", "sys._getframe", "sys.settrace", "sys.setprofile"}: + raise ValueError(f"Protected runtime state in candidate: {module}") if final_language and module.split(".")[0] in {"triton", "cupy", "numba", "aiter", "ctypes", "subprocess"}: raise ValueError(f"Final operator must execute FlyDSL, not {module}") backend_seen |= module == "flydsl" or module.startswith("flydsl.") @@ -114,8 +116,18 @@ def dotted(node): if isinstance(node, ast.Attribute): return dotted(node.value) + "." + node.attr return "" for node in ast.walk(tree): + name = dotted(node) + if name in {"sys.modules", "sys._getframe", "sys.settrace", "sys.setprofile"}: + raise ValueError(f"Protected runtime state in candidate: {name}") + if isinstance(node, ast.Attribute) and node.attr in { + "__dict__", "__globals__", "__builtins__", "__code__", "__closure__", + "__subclasses__", "__getattribute__", "f_globals", "f_locals", "f_back", + }: + raise ValueError(f"Runtime introspection is not allowed: {node.attr}") if not isinstance(node, ast.Call): continue name = dotted(node.func) + if name in {"globals", "locals", "vars", "getattr", "setattr", "delattr"}: + raise ValueError(f"Dynamic runtime introspection is not allowed: {name}") if final_language and (name in {"torch.mm", "torch.bmm", "torch.matmul", "torch.einsum", "torch.softmax", "torch.log_softmax", "torch.layer_norm", "torch.rms_norm"} or name.startswith("torch.nn.functional.")): raise ValueError(f"Library operator shortcut in candidate: {name}") if name in {"eval", "exec", "__import__", "importlib.import_module", "importlib.util.spec_from_file_location"}: diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/README.md b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/README.md index 59d45953e..522ef87eb 100644 --- a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/README.md +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/README.md @@ -77,5 +77,11 @@ Validate the last actual measured output against quantized FP32 GEMM/BF16 cast, then halve both scale tensors outside timing, poison the old output and rerun the same eager callable. Compare the result against the original numerical gate and restore inputs. No captured-graph claim is made for this Event invocation. -The original source imports the older FlyDSL buffer_ops API: qualify it with the -pinned compatible runtime and record that image digest, not an untested image. +The task-local `flydsl_compat` helpers preserve the legacy buffer/vector API +when the installed FlyDSL no longer supplies it. See `flydsl_compat/SOURCE.md` +for the pinned upstream source and retained license. This does not change +workloads, numerical gates, or timing parameters. + +The performance action exposes the existing measured-output verification and +Event metadata as `timed_output_checked` and `device_timing` for the validator. +These fields describe the checks above; they add no timed work. diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/LICENSE b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/LICENSE new file mode 100644 index 000000000..c73e2f2c4 --- /dev/null +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/LICENSE @@ -0,0 +1,17 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + Copyright 2025 FlyDSL Project Contributors + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/SOURCE.md b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/SOURCE.md new file mode 100644 index 000000000..83ea5ad77 --- /dev/null +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/SOURCE.md @@ -0,0 +1,13 @@ +# Pinned FlyDSL compatibility helpers + +Derived from ROCm/FlyDSL commit `28a18d328b4882c999864b2df2f8f9fe3fcc8b47`, +`python/flydsl/expr/{buffer_ops,vector,meta}.py`, under Apache-2.0 (see LICENSE). +The original buffer descriptor, byte offset, masking and cache policy is retained. +Relative dependency imports now target the installed package. Removed memref +pointer extraction uses the current typed iterator and LLVM pointer conversion. +Legacy runtimes continue using their installed original helpers. No runtime +monkeypatching, external repository imports or downloads are performed. + +Current ROCDL load/store operations receive the same cache-policy bits through +their `aux` attribute instead of a removed SSA operand. Vector unwrapping uses +the current signature while retaining the surrounding MLIR location context. diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/__init__.py b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/__init__.py new file mode 100644 index 000000000..1bbff9ffa --- /dev/null +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/__init__.py @@ -0,0 +1,13 @@ +"""Task-local legacy API fallback; never modifies the installed FlyDSL package.""" +try: + import flydsl.expr.buffer_ops as buffer_ops +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.buffer_ops": + raise + from . import buffer_ops +try: + import flydsl.expr.vector as vector +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.vector": + raise + from . import vector diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/buffer_ops.py b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/buffer_ops.py new file mode 100644 index 000000000..ab27dd821 --- /dev/null +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/buffer_ops.py @@ -0,0 +1,603 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""AMD Buffer Load/Store Operations - High-level Python API + +This module provides high-level Python wrappers for AMD CDNA3/CDNA4 buffer operations. +Buffer operations use a scalar base pointer and per-thread offsets for efficient memory access. + +Example: + >>> from flydsl._mlir_helpers import buffer_ops + >>> from flydsl._mlir_helpers import arith + >>> import _mlir.extras.types as T + >>> + >>> # Create buffer resource from memref + >>> rsrc = buffer_ops.create_buffer_resource(A) + >>> + >>> # Compute offset + >>> offset = row * arith.index(4096) + col + >>> + >>> # Buffer load (4xf32) + >>> data = buffer_ops.buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Buffer store + >>> buffer_ops.buffer_store(data, rsrc, offset) +""" + +from typing import Optional, Union + +from flydsl._mlir import ir +from flydsl._mlir.dialects import arith as std_arith +from flydsl._mlir.dialects import llvm, rocdl +from flydsl._mlir.extras import types as T +from flydsl.runtime.device import is_rdna_arch +from .meta import traced_op + + +def _get_buffer_flags(arch=None): + """Get AMD buffer resource descriptor (V#) flags word (bits 127:96). + + Constructs the 32-bit flags field for rocdl.make.buffer.rsrc, following the + same logic as LLVM's AMDGPUToROCDL makeBufferRsrc(): + https://github.com/llvm/llvm-project/blob/main/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp + + Bit layout (common to all architectures): + bits [11:0] - DST_SEL: ignored by raw buffer intrinsics + bits [14:12] - DATA_FORMAT: must be nonzero, 7 = float + bits [18:15] - NUM_FORMAT: must be nonzero, 4 = 32-bit + bit [19] - In nested heap (0) + bit [20] - Behavior on unmap (0 = return 0 / ignore) + bits [22:21] - Index stride for swizzles (0) + bit [23] - Add thread ID (0) + bit [24] - Reserved: must be 1 on RDNA, 0 on CDNA + bits [26:25] - Reserved (0) + bit [27] - Non-volatile (CDNA only, 0) + bits [29:28] - OOB_SELECT (RDNA only): 0=structured, 2=none, 3=check offset + bits [31:30] - Type (must be 0) + + CDNA (gfx9xx): (7 << 12) | (4 << 15) = 0x20070 + RDNA (gfx10+): (7 << 12) | (4 << 15) | (1 << 24) | (2 << 28) = 0x21020070 + - bit 24 set to 1 (required on RDNA) + - OOB_SELECT=2 (no bounds checking, matching LLVM boundsCheck=false) + """ + import os + + if arch is None: + arch = os.environ.get("FLYDSL_GPU_ARCH") + flags = (7 << 12) | (4 << 15) + if is_rdna_arch(arch): + flags |= 1 << 24 # reserved bit, must be 1 on RDNA + flags |= 2 << 28 # OOB_SELECT = 2 (no bounds checking) + return flags + + +__all__ = [ + "create_llvm_ptr", + "get_element_ptr", + "create_buffer_resource", + "create_buffer_resource_from_addr", + "buffer_load", + "buffer_store", + "BufferResourceDescriptor", + "extract_base_index", +] + + +def _unwrap_value(value): + """Recursively unwrap ArithValue or similar wrappers to get the actual MLIR value. + + Handles: + - FlyDSL ArithValue (has ._value) + - flyc DSL Numeric like fx.Int32 (has .ir_value() method) + - flyc ArithValue (is already ir.Value subclass) + """ + # DSL Numeric (Int32, Float32, etc.) — use ir_value() to materialize + if hasattr(value, "ir_value") and not isinstance(value, ir.Value): + return value.ir_value() + max_depth = 10 # Safety limit + depth = 0 + while depth < max_depth and not isinstance(value, ir.Value): + if hasattr(value, "_value"): + value = value._value + elif hasattr(value, "value"): + value = value.value + else: + break + depth += 1 + return value + + +def _create_i32_constant(value: int) -> ir.Value: + """Create i32 constant using standard MLIR arith dialect.""" + i32_type = T.i32() + if value > 0x7FFFFFFF: + value = int(value - 2**32) + attr = ir.IntegerAttr.get(i32_type, value) + op = std_arith.ConstantOp(i32_type, attr) + return _unwrap_value(op.result) + + +def _create_i16_constant(value: int) -> ir.Value: + """Create i16 constant using standard MLIR arith dialect.""" + i16_type = T.i16() + attr = ir.IntegerAttr.get(i16_type, value) + op = std_arith.ConstantOp(i16_type, attr) + return _unwrap_value(op.result) + + +def _create_i64_constant(value: int) -> ir.Value: + """Create i64 constant using standard MLIR arith dialect.""" + i64_type = T.i64() + attr = ir.IntegerAttr.get(i64_type, value) + op = std_arith.ConstantOp(i64_type, attr) + return _unwrap_value(op.result) + + +def create_llvm_ptr(value, address_space: int = 0) -> ir.Value: + """Create an LLVM pointer from an integer or index value.""" + value = _unwrap_value(value) + if isinstance(value.type, ir.IndexType): + i64_type = T.i64() + value = _unwrap_value(std_arith.IndexCastOp(i64_type, value).result) + ptr_type = ir.Type.parse(f"!llvm.ptr<{address_space}>") + return llvm.IntToPtrOp(ptr_type, value).result + + +def extract_base_index(tensor, address_space: int = 1) -> ir.Value: + """Extract the base address of a fly.memref as an index value. + + Inverse of :func:`create_llvm_ptr` (index -> ptr). Useful when ISA + requires a raw pointer instead of a buffer resource descriptor + (e.g. global_atomic_pk_add_bf16 on gfx942). + """ + from flydsl._mlir.dialects import fly as _fly + from flydsl._mlir.dialects import memref as _memref + + raw = _unwrap_value(tensor) + try: + ir.MemRefType(raw.type) + return _memref.extract_aligned_pointer_as_index(raw) + except ValueError: + pass + + # FlyDSL 0.3 uses a typed iterator instead of the removed extract op. + from flydsl.expr import get_iter, to_llvm_ptr + ptr = to_llvm_ptr(get_iter(raw)) + i64_val = llvm.PtrToIntOp(ir.IntegerType.get_signless(64), ptr).result + return _unwrap_value(std_arith.IndexCastOp(ir.IndexType.get(), i64_val).result) + + +def get_element_ptr( + base_ptr, + byte_offset: Union[int, ir.Value, None] = None, + static_byte_offset: int = 0, + elem_type: Optional[ir.Type] = None, + no_wrap_flags=None, +) -> ir.Value: + """Build an LLVM GEP from a base pointer plus byte offsets.""" + _gep_dynamic_index_sentinel = -(2**31) + + base_ptr = _unwrap_value(base_ptr) + if not isinstance(static_byte_offset, int): + raise TypeError(f"static_byte_offset must be int, got {type(static_byte_offset).__name__}") + if elem_type is None: + elem_type = T.i8() + elif callable(elem_type): + elem_type = elem_type() + + if byte_offset is None: + dynamic_indices = [] + raw_constant_indices = [int(static_byte_offset)] + elif isinstance(byte_offset, int): + dynamic_indices = [] + raw_constant_indices = [int(byte_offset) + int(static_byte_offset)] + else: + offset_val = _unwrap_value(byte_offset) + if isinstance(offset_val.type, ir.IndexType): + i64_type = T.i64() + offset_val = _unwrap_value(std_arith.IndexCastOp(i64_type, offset_val).result) + elif not isinstance(offset_val.type, ir.IntegerType): + raise TypeError("byte_offset must be int, index, or integer-typed MLIR value; " f"got {offset_val.type}") + + if static_byte_offset != 0: + static_type = offset_val.type + static_attr = ir.IntegerAttr.get(static_type, int(static_byte_offset)) + static_const = _unwrap_value(std_arith.ConstantOp(static_type, static_attr).result) + offset_val = _unwrap_value(std_arith.AddIOp(offset_val, static_const).result) + + dynamic_indices = [offset_val] + raw_constant_indices = [_gep_dynamic_index_sentinel] + + return llvm.GEPOp( + base_ptr.type, + base_ptr, + dynamic_indices, + raw_constant_indices, + elem_type, + no_wrap_flags, + ).result + + +class BufferResourceDescriptor: + """AMD Buffer Resource Descriptor + + A buffer resource descriptor contains: + - base_pointer: Scalar base pointer (wave-uniform, stored in SGPRs) + - stride: Stride for structured buffers (typically 0 for contiguous) + - num_records: Buffer size in bytes + - flags: Data format and access flags + + The descriptor is stored in a special LLVM pointer type (!llvm.ptr<8>) + """ + + def __init__(self, rsrc: ir.Value): + """Initialize with ROCDL resource descriptor value.""" + self.rsrc = rsrc + + @staticmethod + def from_memref( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + data_format: str = "f32", + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, + ) -> "BufferResourceDescriptor": + """Create buffer resource descriptor from memref. + + Args: + memref_val: Memref value to create descriptor for + stride: Stride in elements (0 for contiguous) + max_size: If True, use max buffer size for flexibility + num_records_bytes: Override buffer size (in BYTES) used by hardware OOB checking. + If provided, this takes precedence over `max_size`. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + data_format: Data format ('f32', 'f16', 'i32', etc.) + + Returns: + BufferResourceDescriptor instance + + Example: + >>> rsrc = BufferResourceDescriptor.from_memref(A) + """ + # Extract raw pointer from fly.memref. + raw_val = _unwrap_value(memref_val) + from flydsl._mlir.dialects import fly as _fly + + # Preserve byte offsets and descriptor policy with the current pointer API. + from flydsl.expr import get_iter, to_llvm_ptr + base_ptr = to_llvm_ptr(get_iter(raw_val)) + if base_byte_offset is not None: + base_ptr = get_element_ptr(base_ptr, byte_offset=base_byte_offset) + + # Create buffer resource descriptor + flags_val = _get_buffer_flags() + flags = _create_i32_constant(flags_val) + stride_val = _create_i16_constant(stride) + + def _num_records_from_memref_type() -> Optional[int]: + """Best-effort: derive logical buffer size (in bytes) from static memref type.""" + try: + mt = ir.MemRefType(_unwrap_value(memref_val).type) + shape = list(mt.shape) + if any(int(d) < 0 for d in shape): + return None + # Compute element size in bytes (scalar element type). + elem_t = mt.element_type + elem_bits = getattr(elem_t, "width", None) + if elem_bits is None: + return None + elem_bytes = int(elem_bits) // 8 + if elem_bytes <= 0: + return None + num_elems = 1 + for d in shape: + num_elems *= int(d) + return int(num_elems) * int(elem_bytes) + except Exception: + return None + + if num_records_bytes is not None: + # Caller-provided size in BYTES (preferred for exact hardware OOB behavior). + if isinstance(num_records_bytes, int): + nbytes = int(num_records_bytes) + if nbytes <= 0: + nbytes = 0 + # Descriptor uses i32 bytes; clamp to the max representable. + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(nbytes) + else: + v = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(v.type, ir.IntegerType) or v.type.width != 64: + if isinstance(v.type, ir.IndexType): + op = std_arith.IndexCastOp(i64_type, v) + else: + op = std_arith.ExtSIOp(i64_type, v) + v = _unwrap_value(op.result) + num_records = v + elif max_size: + # Use max for flexibility (hardware will check actual bounds) + # Note: FlyDSL's rocdl.make.buffer.rsrc requires i32, not i64 + num_records = _create_i64_constant(0xFFFFFFFF) # FALLBACK_MAX_SIZE + else: + # Use the logical memref size (in bytes) for hardware OOB checking. + nbytes = _num_records_from_memref_type() + if nbytes is None: + # Fall back to max-size if we can't infer statically. + num_records = _create_i64_constant(0xFFFFFFFF) + else: + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(int(nbytes)) + + # Create resource descriptor (returns !llvm.ptr<8>) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + rsrc = rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride_val, num_records, flags).result + + return BufferResourceDescriptor(rsrc) + + +def create_buffer_resource_from_addr( + addr_i64: ir.Value, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from a raw i64 device address. + + Useful when working with runtime pointer arrays (e.g. IPC-mapped addresses + or device-side pointer tables) where no fly.memref is available. + The full address is encoded as the buffer base; callers should pass + byte offset 0 to buffer_load / buffer_store. + + Args: + addr_i64: Raw 64-bit device address (i64 MLIR value). + num_records_bytes: Optional buffer size in bytes for hardware OOB checking. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>). + + Example: + >>> rsrc = create_buffer_resource_from_addr(raw_addr_i64) + >>> data = buffer_load(rsrc, i32_zero, vec_width=4, dtype=T.i32) + """ + addr_i64 = _unwrap_value(addr_i64) + ptr_type = ir.Type.parse("!llvm.ptr") + base_ptr = llvm.IntToPtrOp(ptr_type, addr_i64).result + flags = _create_i32_constant(_get_buffer_flags()) + stride = _create_i16_constant(0) + if num_records_bytes is None: + num_records = _create_i64_constant(0xFFFFFFFF) + elif isinstance(num_records_bytes, int): + nbytes = max(0, min(int(num_records_bytes), 0xFFFFFFFF)) + num_records = _create_i64_constant(nbytes) + else: + num_records = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(num_records.type, ir.IntegerType) or num_records.type.width != 64: + if isinstance(num_records.type, ir.IndexType): + num_records = _unwrap_value(std_arith.IndexCastOp(i64_type, num_records).result) + else: + num_records = _unwrap_value(std_arith.ExtSIOp(i64_type, num_records).result) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + return rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride, num_records, flags).result + + +@traced_op +def create_buffer_resource( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from memref. + + This is a simplified wrapper around BufferResourceDescriptor.from_memref() + that returns the raw ROCDL resource value. + + Args: + memref_val: Memref value + stride: Buffer stride (0 for contiguous) + max_size: Use maximum buffer size + num_records_bytes: Override buffer size in bytes. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>) + + Example: + >>> rsrc = create_buffer_resource(A) + >>> data = buffer_load(rsrc, offset) + """ + desc = BufferResourceDescriptor.from_memref( + memref_val, + stride, + max_size, + num_records_bytes=num_records_bytes, + base_byte_offset=base_byte_offset, + ) + return desc.rsrc + + +@traced_op +def buffer_load( + rsrc: ir.Value, + offset: ir.Value, + vec_width: int = 4, + dtype=None, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + soffset_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """AMD buffer load operation. + + Load data from global memory using buffer descriptor and offset. + Uses hardware-level bounds checking and vectorization. + + Args: + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + vec_width: Vector width (1, 2, or 4) + dtype: Element data type (None for f32, or ir.F32Type, etc.) + mask: Optional mask for predicated load (i1 type) + cache_modifier: Cache control flags (0 for default) + soffset_bytes: Optional scalar offset (in BYTES) added by the buffer instruction (soffset). + Use this to fold small constant deltas into the instruction instead of emitting + extra VGPR address arithmetic. + + Returns: + Loaded data (scalar or vector depending on vec_width) + + Example: + >>> # Load 4xf32 + >>> data = buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Load with mask + >>> data = buffer_load(rsrc, offset, vec_width=4, mask=valid) + """ + # Default dtype to f32 + if dtype is None: + dtype = T.f32() + # Accept DSL Numeric class (e.g. fx.Int32) as dtype: unwrap to ir.Type + elif hasattr(dtype, "ir_type"): + dtype = dtype.ir_type + + # Unwrap offset first (accept Python ints and DSL Numeric values). + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: Buffer load offset is in BYTES, not elements! + # For vec4xf32, each element is 4 bytes, so multiply offset by 4 + element_bytes = dtype.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create vector type + if vec_width == 1: + result_type = dtype + else: + result_type = ir.VectorType.get([vec_width], dtype) + + # Create instruction offset and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(soffset_bytes) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer load + load_op = rocdl.RawPtrBufferLoadOp( + result_type, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) + + return load_op.result + + +@traced_op +def buffer_store( + data: ir.Value, + rsrc: ir.Value, + offset: ir.Value, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + *, + soffset_bytes: Optional[Union[int, ir.Value]] = None, + offset_is_bytes: bool = False, +): + """AMD buffer store operation. + + Store data to global memory using buffer descriptor and offset. + + Args: + data: Data to store (scalar or vector) + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + mask: Optional mask for predicated store (i1 type) + cache_modifier: Cache control flags (0 for default) + + Example: + >>> buffer_store(data, rsrc, offset) + >>> + >>> # Store with mask + >>> buffer_store(data, rsrc, offset, mask=valid) + """ + # Unwrap all inputs (accept DSL Numeric values via ir_value()) + if hasattr(data, "ir_value"): + data = data.ir_value() + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + data = _unwrap_value(data) + rsrc = _unwrap_value(rsrc) + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: RawPtrBufferStoreOp offset is in BYTES. + # For backward compat, `buffer_store()` accepts element offsets by default + # and scales them to bytes. Set `offset_is_bytes=True` to skip scaling. + if not offset_is_bytes: + # Get element size from data type + data_type = data.type + if hasattr(data_type, "element_type"): # Vector type + element_type = data_type.element_type + else: # Scalar type + element_type = data_type + element_bytes = element_type.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create instruction offset (soffset) and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(int(soffset_bytes)) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer store + rocdl.RawPtrBufferStoreOp( + data, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/meta.py b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/meta.py new file mode 100644 index 000000000..23e4c849a --- /dev/null +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/meta.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +import inspect +from functools import wraps + +from flydsl._mlir import ir + + +def _to_raw_value(obj): + if isinstance(obj, ir.Value): + return obj + if isinstance(obj, type): + return obj + if hasattr(obj, "__extract_to_ir_values__"): + values = obj.__extract_to_ir_values__() + if len(values) != 1: + raise ValueError(f"Primitive function expects 1 value, got {len(values)}") + return values[0] + if isinstance(obj, tuple): + return tuple(_to_raw_value(e) for e in obj) + if isinstance(obj, list): + return [_to_raw_value(e) for e in obj] + return obj + + +def _flatten_args(args, kwargs): + new_args = tuple(_to_raw_value(a) for a in args) + new_kwargs = {k: _to_raw_value(v) if k not in ("loc", "ip") else v for k, v in kwargs.items()} + return new_args, new_kwargs + + +def _caller_location(depth=1): + """Build an MLIR Location from the Python call-site *depth* frames up.""" + frame = inspect.currentframe() + for _ in range(depth + 1): + if frame is not None: + frame = frame.f_back + if frame is None: + return ir.Location.unknown() + + info = inspect.getframeinfo(frame) + pos = getattr(info, "positions", None) + line = pos.lineno if pos is not None else info.lineno + col = (pos.col_offset or 0) if pos is not None else 0 + file_loc = ir.Location.file(info.filename, line, col) + + if info.code_context: + label = " ".join(ln.strip() for ln in info.code_context) + else: + label = info.function + return ir.Location.name(label, childLoc=file_loc) + + +def traced_op(op): + @wraps(op) + def wrapper(*args, **kwargs): + loc = kwargs.pop("loc", None) + if loc is None: + loc = _caller_location(depth=1) + args, kwargs = _flatten_args(args, kwargs) + with loc: + return op(*args, **kwargs) + + return wrapper diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/vector.py b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/vector.py new file mode 100644 index 000000000..29f8289c2 --- /dev/null +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/flydsl_compat/vector.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""Vector dialect helpers and re-exports. + +The ``Vector`` class itself lives in ``typing.py`` alongside other builtin +DSL types. This module re-exports it for convenience and provides thin +wrappers around upstream ``_mlir.dialects.vector`` ops. +""" + +from __future__ import annotations + +from flydsl._mlir import ir +from flydsl._mlir.dialects import vector as _vector + +# Re-export upstream dialect for ``from flydsl.expr import vector; vector.broadcast(...)`` +from flydsl._mlir.dialects.vector import * # noqa: F401,F403,E402 +from .meta import traced_op + +# Re-export Vector and friends so ``from flydsl.expr.vector import Vector`` works +from flydsl.expr.typing import ReductionOp, Vector, empty_like, full, full_like, ones_like, zeros_like # noqa: F401 + +# ═══════════════════════════════════════════════════════════════════════ +# Dialect helper wrappers (legacy, will be deprecated) +# Prefer using Vector methods or _mlir.dialects.vector directly. +# ═══════════════════════════════════════════════════════════════════════ + + +@traced_op +def from_elements(*args, loc=None, ip=None, **kwargs): + """Construct a vector from scalar elements, auto-unwrapping ArithValue wrappers.""" + from flydsl.expr import arith as _arith_ext + + if len(args) >= 2: + args = list(args) + elems = args[1] + if isinstance(elems, (list, tuple)): + args[1] = [_arith_ext.unwrap(v) for v in elems] + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + +@traced_op +def store(value, memref, indices, *, loc=None, ip=None, **kwargs): + """Vector store wrapper that accepts ArithValue/wrappers for value/indices.""" + from flydsl.expr import arith as _arith_ext + + return _vector.store( + _arith_ext.unwrap(value), + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + **kwargs, + ) + + +# ----------------------------------------------------------------------------- +# Thin wrappers for common op classes that otherwise require `.result` access. +# ----------------------------------------------------------------------------- + + +@traced_op +def extract(vector, static_position=None, dynamic_position=None, *, loc=None, ip=None): + """Wrapper around `vector.ExtractOp(...).result`. + + When only ``dynamic_position`` is supplied (without explicit + ``static_position``), each dynamic index needs a corresponding + ``kDynamic`` sentinel in the static attribute so the ODS builder + pairs them correctly. This wrapper fills in the sentinels + automatically. + """ + from flydsl.expr import arith as _arith_ext + + if static_position is None: + static_position = [] + if dynamic_position is None: + dynamic_position = [] + dynamic_position = [_arith_ext.unwrap(i, index=True) for i in dynamic_position] + + n_static = len(static_position) + n_dynamic = len(dynamic_position) + if n_dynamic > 0 and n_static < n_dynamic: + kDynamic = ir.ShapedType.get_dynamic_size() + static_position = list(static_position) + [kDynamic] * (n_dynamic - n_static) + + return _vector.ExtractOp( + _arith_ext.unwrap(vector), + static_position=static_position, + dynamic_position=dynamic_position, + loc=loc, + ip=ip, + ).result + + +@traced_op +def load_op(result_type, memref, indices, *, loc=None, ip=None): + """Wrapper around `vector.LoadOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.LoadOp( + result_type, + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + ).result + + +@traced_op +def bitcast(result_type, source, *, loc=None, ip=None): + """Wrapper around `vector.BitCastOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.BitCastOp( + result_type, + _arith_ext.unwrap(source), + loc=loc, + ip=ip, + ).result diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/kernel.py b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/kernel.py index 59def19e9..27d12bd46 100644 --- a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/kernel.py +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/kernel.py @@ -32,7 +32,8 @@ from flydsl._mlir.dialects.arith import CmpIPredicate from flydsl.compiler.kernel_function import CompilationContext from flydsl.expr import arith as _arith -from flydsl.expr import buffer_ops, const_expr, gpu, math, range_constexpr, rocdl +from flydsl.expr import const_expr, gpu, math, range_constexpr, rocdl +from flydsl_compat import buffer_ops, vector from flydsl.expr.typing import T from flydsl.runtime.device import get_rocm_arch as get_hip_arch from flydsl.utils.smem_allocator import SmemAllocator, SmemPtr @@ -1760,7 +1761,7 @@ def load_b_pack(base_k, ki_step, ni): return load_b_pack_k32( buffer_ops, fx.arith, - fx.vector, + vector, arg_b=arg_b, b_rsrc=b_rsrc, layout_b=layout_b, @@ -1838,7 +1839,7 @@ def load_b_packs_k64(base_k, ku: int, ni: int): vec_elems = 16 if elem_bytes == 1 else 8 b16 = _buffer_load_vec( buffer_ops, - fx.vector, + vector, b_rsrc, idx_pack, elem_type=_elem_type(), @@ -1919,7 +1920,7 @@ def lds_load_packs_k64(curr_row_a_lds, col_base, lds_buffer): def load_a_16(idx_elem): return buffer_copy_gmem16_dwordx4( buffer_ops, - fx.vector, + vector, elem_type=_elem_type(), idx_i32=idx_elem, rsrc=a_rsrc, @@ -2460,7 +2461,7 @@ def store_pair(*, row_local, row, row_ctx, col_pair0, col_g0, frag): mfma_epilog( use_cshuffle=True, arith=fx.arith, - vector=fx.vector, + vector=vector, gpu=gpu, range_constexpr=range_constexpr, tile_m=tile_m, diff --git a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/scripts/task_actions.py b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/scripts/task_actions.py index 9fcb240d6..a19eed6d9 100644 --- a/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/scripts/task_actions.py +++ b/tasks/torch2flydsl/gemm_a8w8_bpreshuffle_kernel/scripts/task_actions.py @@ -15,4 +15,12 @@ def check(h): return sorted(observed) def performance(h): - return h.arena_benchmark(warmup=10, iters=100, verbose=True) + rows = h.arena_benchmark(warmup=10, iters=100, verbose=True) + for row in rows: + # Preserve the measured values and expose the existing post-timing + # checks through the validator's event-timing evidence contract. + row["device_timing"] = { + key: value for key, value in row.items() if key.startswith("benchmark_") + } + row["timed_output_checked"] = row.get("timed_output_correctness") == "PASS" + return rows diff --git a/tasks/torch2flydsl/hgemm_kernel/README.md b/tasks/torch2flydsl/hgemm_kernel/README.md index 9f8bdc198..be8e99f69 100644 --- a/tasks/torch2flydsl/hgemm_kernel/README.md +++ b/tasks/torch2flydsl/hgemm_kernel/README.md @@ -61,17 +61,12 @@ re-invokes the same eager callable and is not captured graph replay. Graph timin validates the actual captured replay. Original shapes, seeds, tolerance, warmup, sample counts and Graph/Event selection remain unchanged. -Runtime qualification: the initial implementation uses the legacy FlyDSL -`expr.buffer_ops` and `expr.vector` APIs. It passed the full task validator on -MI355X with the repository's pinned SGLang 0.5.14 / FlyDSL 0.2.2 runtime -(2026-09-15, job 139392 on an available-memory MI355X allocation; all five -correctness and performance cases, including -measured Event outputs and eager re-invocation). The tested SGLang 0.5.19 / -FlyDSL 0.3.2 runtime removes these APIs and fails before kernel execution. -Select the qualified runtime through the run-level Docker image setting; -an unchanged source under 0.5.19 is not qualified. Candidate and baseline must -use the same runtime and timing method. This initial-task validation does not -certify a subsequently modified candidate. +The initial implementation uses the legacy FlyDSL buffer/vector interface. +The bundled compatibility package supplies the removed helpers on current +FlyDSL runtimes; see its attribution below. Select the runtime through the +run-level Docker image setting, and keep baseline and candidate on the same +image and timing method. Qualification binds a specific task source and image +and does not certify a subsequently modified candidate. Only the entrypoints listed in config.yaml are required interfaces. The primary @@ -90,3 +85,18 @@ also forbidden. Ordinary Python utilities, PyTorch allocation/layout operations, and the task's bundled `kernels/` helpers remain available under the existing numerical and timing contract. Baseline checks retain their declared initial backend; the final candidate must use FlyDSL. + +The task-local `flydsl_compat` helpers preserve the legacy buffer/vector API +when the installed FlyDSL no longer supplies it. See `flydsl_compat/SOURCE.md` +for the pinned upstream source and retained license. Workloads, numerical gates, +and timing parameters are unchanged. + +The performance action exposes the existing measured-output verification and +Event metadata as `timed_output_checked` and `device_timing` for the validator. +These fields add no timed work and preserve the protected benchmark settings. + +Candidate imports must not mutate shared dependency objects such as `torch.matmul`, +including through import aliases. The task rejects known module/frame introspection +routes before importing candidate code; the local compiled-kernel `_cf` cache lookup +remains allowed. These checks supplement immutable harness files and GPU output +checks; they are not a security sandbox for arbitrary Python. diff --git a/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/LICENSE b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/LICENSE new file mode 100644 index 000000000..c73e2f2c4 --- /dev/null +++ b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/LICENSE @@ -0,0 +1,17 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + Copyright 2025 FlyDSL Project Contributors + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/SOURCE.md b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/SOURCE.md new file mode 100644 index 000000000..83ea5ad77 --- /dev/null +++ b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/SOURCE.md @@ -0,0 +1,13 @@ +# Pinned FlyDSL compatibility helpers + +Derived from ROCm/FlyDSL commit `28a18d328b4882c999864b2df2f8f9fe3fcc8b47`, +`python/flydsl/expr/{buffer_ops,vector,meta}.py`, under Apache-2.0 (see LICENSE). +The original buffer descriptor, byte offset, masking and cache policy is retained. +Relative dependency imports now target the installed package. Removed memref +pointer extraction uses the current typed iterator and LLVM pointer conversion. +Legacy runtimes continue using their installed original helpers. No runtime +monkeypatching, external repository imports or downloads are performed. + +Current ROCDL load/store operations receive the same cache-policy bits through +their `aux` attribute instead of a removed SSA operand. Vector unwrapping uses +the current signature while retaining the surrounding MLIR location context. diff --git a/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/__init__.py b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/__init__.py new file mode 100644 index 000000000..1bbff9ffa --- /dev/null +++ b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/__init__.py @@ -0,0 +1,13 @@ +"""Task-local legacy API fallback; never modifies the installed FlyDSL package.""" +try: + import flydsl.expr.buffer_ops as buffer_ops +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.buffer_ops": + raise + from . import buffer_ops +try: + import flydsl.expr.vector as vector +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.vector": + raise + from . import vector diff --git a/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/buffer_ops.py b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/buffer_ops.py new file mode 100644 index 000000000..ab27dd821 --- /dev/null +++ b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/buffer_ops.py @@ -0,0 +1,603 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""AMD Buffer Load/Store Operations - High-level Python API + +This module provides high-level Python wrappers for AMD CDNA3/CDNA4 buffer operations. +Buffer operations use a scalar base pointer and per-thread offsets for efficient memory access. + +Example: + >>> from flydsl._mlir_helpers import buffer_ops + >>> from flydsl._mlir_helpers import arith + >>> import _mlir.extras.types as T + >>> + >>> # Create buffer resource from memref + >>> rsrc = buffer_ops.create_buffer_resource(A) + >>> + >>> # Compute offset + >>> offset = row * arith.index(4096) + col + >>> + >>> # Buffer load (4xf32) + >>> data = buffer_ops.buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Buffer store + >>> buffer_ops.buffer_store(data, rsrc, offset) +""" + +from typing import Optional, Union + +from flydsl._mlir import ir +from flydsl._mlir.dialects import arith as std_arith +from flydsl._mlir.dialects import llvm, rocdl +from flydsl._mlir.extras import types as T +from flydsl.runtime.device import is_rdna_arch +from .meta import traced_op + + +def _get_buffer_flags(arch=None): + """Get AMD buffer resource descriptor (V#) flags word (bits 127:96). + + Constructs the 32-bit flags field for rocdl.make.buffer.rsrc, following the + same logic as LLVM's AMDGPUToROCDL makeBufferRsrc(): + https://github.com/llvm/llvm-project/blob/main/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp + + Bit layout (common to all architectures): + bits [11:0] - DST_SEL: ignored by raw buffer intrinsics + bits [14:12] - DATA_FORMAT: must be nonzero, 7 = float + bits [18:15] - NUM_FORMAT: must be nonzero, 4 = 32-bit + bit [19] - In nested heap (0) + bit [20] - Behavior on unmap (0 = return 0 / ignore) + bits [22:21] - Index stride for swizzles (0) + bit [23] - Add thread ID (0) + bit [24] - Reserved: must be 1 on RDNA, 0 on CDNA + bits [26:25] - Reserved (0) + bit [27] - Non-volatile (CDNA only, 0) + bits [29:28] - OOB_SELECT (RDNA only): 0=structured, 2=none, 3=check offset + bits [31:30] - Type (must be 0) + + CDNA (gfx9xx): (7 << 12) | (4 << 15) = 0x20070 + RDNA (gfx10+): (7 << 12) | (4 << 15) | (1 << 24) | (2 << 28) = 0x21020070 + - bit 24 set to 1 (required on RDNA) + - OOB_SELECT=2 (no bounds checking, matching LLVM boundsCheck=false) + """ + import os + + if arch is None: + arch = os.environ.get("FLYDSL_GPU_ARCH") + flags = (7 << 12) | (4 << 15) + if is_rdna_arch(arch): + flags |= 1 << 24 # reserved bit, must be 1 on RDNA + flags |= 2 << 28 # OOB_SELECT = 2 (no bounds checking) + return flags + + +__all__ = [ + "create_llvm_ptr", + "get_element_ptr", + "create_buffer_resource", + "create_buffer_resource_from_addr", + "buffer_load", + "buffer_store", + "BufferResourceDescriptor", + "extract_base_index", +] + + +def _unwrap_value(value): + """Recursively unwrap ArithValue or similar wrappers to get the actual MLIR value. + + Handles: + - FlyDSL ArithValue (has ._value) + - flyc DSL Numeric like fx.Int32 (has .ir_value() method) + - flyc ArithValue (is already ir.Value subclass) + """ + # DSL Numeric (Int32, Float32, etc.) — use ir_value() to materialize + if hasattr(value, "ir_value") and not isinstance(value, ir.Value): + return value.ir_value() + max_depth = 10 # Safety limit + depth = 0 + while depth < max_depth and not isinstance(value, ir.Value): + if hasattr(value, "_value"): + value = value._value + elif hasattr(value, "value"): + value = value.value + else: + break + depth += 1 + return value + + +def _create_i32_constant(value: int) -> ir.Value: + """Create i32 constant using standard MLIR arith dialect.""" + i32_type = T.i32() + if value > 0x7FFFFFFF: + value = int(value - 2**32) + attr = ir.IntegerAttr.get(i32_type, value) + op = std_arith.ConstantOp(i32_type, attr) + return _unwrap_value(op.result) + + +def _create_i16_constant(value: int) -> ir.Value: + """Create i16 constant using standard MLIR arith dialect.""" + i16_type = T.i16() + attr = ir.IntegerAttr.get(i16_type, value) + op = std_arith.ConstantOp(i16_type, attr) + return _unwrap_value(op.result) + + +def _create_i64_constant(value: int) -> ir.Value: + """Create i64 constant using standard MLIR arith dialect.""" + i64_type = T.i64() + attr = ir.IntegerAttr.get(i64_type, value) + op = std_arith.ConstantOp(i64_type, attr) + return _unwrap_value(op.result) + + +def create_llvm_ptr(value, address_space: int = 0) -> ir.Value: + """Create an LLVM pointer from an integer or index value.""" + value = _unwrap_value(value) + if isinstance(value.type, ir.IndexType): + i64_type = T.i64() + value = _unwrap_value(std_arith.IndexCastOp(i64_type, value).result) + ptr_type = ir.Type.parse(f"!llvm.ptr<{address_space}>") + return llvm.IntToPtrOp(ptr_type, value).result + + +def extract_base_index(tensor, address_space: int = 1) -> ir.Value: + """Extract the base address of a fly.memref as an index value. + + Inverse of :func:`create_llvm_ptr` (index -> ptr). Useful when ISA + requires a raw pointer instead of a buffer resource descriptor + (e.g. global_atomic_pk_add_bf16 on gfx942). + """ + from flydsl._mlir.dialects import fly as _fly + from flydsl._mlir.dialects import memref as _memref + + raw = _unwrap_value(tensor) + try: + ir.MemRefType(raw.type) + return _memref.extract_aligned_pointer_as_index(raw) + except ValueError: + pass + + # FlyDSL 0.3 uses a typed iterator instead of the removed extract op. + from flydsl.expr import get_iter, to_llvm_ptr + ptr = to_llvm_ptr(get_iter(raw)) + i64_val = llvm.PtrToIntOp(ir.IntegerType.get_signless(64), ptr).result + return _unwrap_value(std_arith.IndexCastOp(ir.IndexType.get(), i64_val).result) + + +def get_element_ptr( + base_ptr, + byte_offset: Union[int, ir.Value, None] = None, + static_byte_offset: int = 0, + elem_type: Optional[ir.Type] = None, + no_wrap_flags=None, +) -> ir.Value: + """Build an LLVM GEP from a base pointer plus byte offsets.""" + _gep_dynamic_index_sentinel = -(2**31) + + base_ptr = _unwrap_value(base_ptr) + if not isinstance(static_byte_offset, int): + raise TypeError(f"static_byte_offset must be int, got {type(static_byte_offset).__name__}") + if elem_type is None: + elem_type = T.i8() + elif callable(elem_type): + elem_type = elem_type() + + if byte_offset is None: + dynamic_indices = [] + raw_constant_indices = [int(static_byte_offset)] + elif isinstance(byte_offset, int): + dynamic_indices = [] + raw_constant_indices = [int(byte_offset) + int(static_byte_offset)] + else: + offset_val = _unwrap_value(byte_offset) + if isinstance(offset_val.type, ir.IndexType): + i64_type = T.i64() + offset_val = _unwrap_value(std_arith.IndexCastOp(i64_type, offset_val).result) + elif not isinstance(offset_val.type, ir.IntegerType): + raise TypeError("byte_offset must be int, index, or integer-typed MLIR value; " f"got {offset_val.type}") + + if static_byte_offset != 0: + static_type = offset_val.type + static_attr = ir.IntegerAttr.get(static_type, int(static_byte_offset)) + static_const = _unwrap_value(std_arith.ConstantOp(static_type, static_attr).result) + offset_val = _unwrap_value(std_arith.AddIOp(offset_val, static_const).result) + + dynamic_indices = [offset_val] + raw_constant_indices = [_gep_dynamic_index_sentinel] + + return llvm.GEPOp( + base_ptr.type, + base_ptr, + dynamic_indices, + raw_constant_indices, + elem_type, + no_wrap_flags, + ).result + + +class BufferResourceDescriptor: + """AMD Buffer Resource Descriptor + + A buffer resource descriptor contains: + - base_pointer: Scalar base pointer (wave-uniform, stored in SGPRs) + - stride: Stride for structured buffers (typically 0 for contiguous) + - num_records: Buffer size in bytes + - flags: Data format and access flags + + The descriptor is stored in a special LLVM pointer type (!llvm.ptr<8>) + """ + + def __init__(self, rsrc: ir.Value): + """Initialize with ROCDL resource descriptor value.""" + self.rsrc = rsrc + + @staticmethod + def from_memref( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + data_format: str = "f32", + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, + ) -> "BufferResourceDescriptor": + """Create buffer resource descriptor from memref. + + Args: + memref_val: Memref value to create descriptor for + stride: Stride in elements (0 for contiguous) + max_size: If True, use max buffer size for flexibility + num_records_bytes: Override buffer size (in BYTES) used by hardware OOB checking. + If provided, this takes precedence over `max_size`. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + data_format: Data format ('f32', 'f16', 'i32', etc.) + + Returns: + BufferResourceDescriptor instance + + Example: + >>> rsrc = BufferResourceDescriptor.from_memref(A) + """ + # Extract raw pointer from fly.memref. + raw_val = _unwrap_value(memref_val) + from flydsl._mlir.dialects import fly as _fly + + # Preserve byte offsets and descriptor policy with the current pointer API. + from flydsl.expr import get_iter, to_llvm_ptr + base_ptr = to_llvm_ptr(get_iter(raw_val)) + if base_byte_offset is not None: + base_ptr = get_element_ptr(base_ptr, byte_offset=base_byte_offset) + + # Create buffer resource descriptor + flags_val = _get_buffer_flags() + flags = _create_i32_constant(flags_val) + stride_val = _create_i16_constant(stride) + + def _num_records_from_memref_type() -> Optional[int]: + """Best-effort: derive logical buffer size (in bytes) from static memref type.""" + try: + mt = ir.MemRefType(_unwrap_value(memref_val).type) + shape = list(mt.shape) + if any(int(d) < 0 for d in shape): + return None + # Compute element size in bytes (scalar element type). + elem_t = mt.element_type + elem_bits = getattr(elem_t, "width", None) + if elem_bits is None: + return None + elem_bytes = int(elem_bits) // 8 + if elem_bytes <= 0: + return None + num_elems = 1 + for d in shape: + num_elems *= int(d) + return int(num_elems) * int(elem_bytes) + except Exception: + return None + + if num_records_bytes is not None: + # Caller-provided size in BYTES (preferred for exact hardware OOB behavior). + if isinstance(num_records_bytes, int): + nbytes = int(num_records_bytes) + if nbytes <= 0: + nbytes = 0 + # Descriptor uses i32 bytes; clamp to the max representable. + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(nbytes) + else: + v = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(v.type, ir.IntegerType) or v.type.width != 64: + if isinstance(v.type, ir.IndexType): + op = std_arith.IndexCastOp(i64_type, v) + else: + op = std_arith.ExtSIOp(i64_type, v) + v = _unwrap_value(op.result) + num_records = v + elif max_size: + # Use max for flexibility (hardware will check actual bounds) + # Note: FlyDSL's rocdl.make.buffer.rsrc requires i32, not i64 + num_records = _create_i64_constant(0xFFFFFFFF) # FALLBACK_MAX_SIZE + else: + # Use the logical memref size (in bytes) for hardware OOB checking. + nbytes = _num_records_from_memref_type() + if nbytes is None: + # Fall back to max-size if we can't infer statically. + num_records = _create_i64_constant(0xFFFFFFFF) + else: + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(int(nbytes)) + + # Create resource descriptor (returns !llvm.ptr<8>) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + rsrc = rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride_val, num_records, flags).result + + return BufferResourceDescriptor(rsrc) + + +def create_buffer_resource_from_addr( + addr_i64: ir.Value, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from a raw i64 device address. + + Useful when working with runtime pointer arrays (e.g. IPC-mapped addresses + or device-side pointer tables) where no fly.memref is available. + The full address is encoded as the buffer base; callers should pass + byte offset 0 to buffer_load / buffer_store. + + Args: + addr_i64: Raw 64-bit device address (i64 MLIR value). + num_records_bytes: Optional buffer size in bytes for hardware OOB checking. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>). + + Example: + >>> rsrc = create_buffer_resource_from_addr(raw_addr_i64) + >>> data = buffer_load(rsrc, i32_zero, vec_width=4, dtype=T.i32) + """ + addr_i64 = _unwrap_value(addr_i64) + ptr_type = ir.Type.parse("!llvm.ptr") + base_ptr = llvm.IntToPtrOp(ptr_type, addr_i64).result + flags = _create_i32_constant(_get_buffer_flags()) + stride = _create_i16_constant(0) + if num_records_bytes is None: + num_records = _create_i64_constant(0xFFFFFFFF) + elif isinstance(num_records_bytes, int): + nbytes = max(0, min(int(num_records_bytes), 0xFFFFFFFF)) + num_records = _create_i64_constant(nbytes) + else: + num_records = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(num_records.type, ir.IntegerType) or num_records.type.width != 64: + if isinstance(num_records.type, ir.IndexType): + num_records = _unwrap_value(std_arith.IndexCastOp(i64_type, num_records).result) + else: + num_records = _unwrap_value(std_arith.ExtSIOp(i64_type, num_records).result) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + return rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride, num_records, flags).result + + +@traced_op +def create_buffer_resource( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from memref. + + This is a simplified wrapper around BufferResourceDescriptor.from_memref() + that returns the raw ROCDL resource value. + + Args: + memref_val: Memref value + stride: Buffer stride (0 for contiguous) + max_size: Use maximum buffer size + num_records_bytes: Override buffer size in bytes. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>) + + Example: + >>> rsrc = create_buffer_resource(A) + >>> data = buffer_load(rsrc, offset) + """ + desc = BufferResourceDescriptor.from_memref( + memref_val, + stride, + max_size, + num_records_bytes=num_records_bytes, + base_byte_offset=base_byte_offset, + ) + return desc.rsrc + + +@traced_op +def buffer_load( + rsrc: ir.Value, + offset: ir.Value, + vec_width: int = 4, + dtype=None, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + soffset_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """AMD buffer load operation. + + Load data from global memory using buffer descriptor and offset. + Uses hardware-level bounds checking and vectorization. + + Args: + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + vec_width: Vector width (1, 2, or 4) + dtype: Element data type (None for f32, or ir.F32Type, etc.) + mask: Optional mask for predicated load (i1 type) + cache_modifier: Cache control flags (0 for default) + soffset_bytes: Optional scalar offset (in BYTES) added by the buffer instruction (soffset). + Use this to fold small constant deltas into the instruction instead of emitting + extra VGPR address arithmetic. + + Returns: + Loaded data (scalar or vector depending on vec_width) + + Example: + >>> # Load 4xf32 + >>> data = buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Load with mask + >>> data = buffer_load(rsrc, offset, vec_width=4, mask=valid) + """ + # Default dtype to f32 + if dtype is None: + dtype = T.f32() + # Accept DSL Numeric class (e.g. fx.Int32) as dtype: unwrap to ir.Type + elif hasattr(dtype, "ir_type"): + dtype = dtype.ir_type + + # Unwrap offset first (accept Python ints and DSL Numeric values). + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: Buffer load offset is in BYTES, not elements! + # For vec4xf32, each element is 4 bytes, so multiply offset by 4 + element_bytes = dtype.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create vector type + if vec_width == 1: + result_type = dtype + else: + result_type = ir.VectorType.get([vec_width], dtype) + + # Create instruction offset and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(soffset_bytes) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer load + load_op = rocdl.RawPtrBufferLoadOp( + result_type, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) + + return load_op.result + + +@traced_op +def buffer_store( + data: ir.Value, + rsrc: ir.Value, + offset: ir.Value, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + *, + soffset_bytes: Optional[Union[int, ir.Value]] = None, + offset_is_bytes: bool = False, +): + """AMD buffer store operation. + + Store data to global memory using buffer descriptor and offset. + + Args: + data: Data to store (scalar or vector) + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + mask: Optional mask for predicated store (i1 type) + cache_modifier: Cache control flags (0 for default) + + Example: + >>> buffer_store(data, rsrc, offset) + >>> + >>> # Store with mask + >>> buffer_store(data, rsrc, offset, mask=valid) + """ + # Unwrap all inputs (accept DSL Numeric values via ir_value()) + if hasattr(data, "ir_value"): + data = data.ir_value() + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + data = _unwrap_value(data) + rsrc = _unwrap_value(rsrc) + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: RawPtrBufferStoreOp offset is in BYTES. + # For backward compat, `buffer_store()` accepts element offsets by default + # and scales them to bytes. Set `offset_is_bytes=True` to skip scaling. + if not offset_is_bytes: + # Get element size from data type + data_type = data.type + if hasattr(data_type, "element_type"): # Vector type + element_type = data_type.element_type + else: # Scalar type + element_type = data_type + element_bytes = element_type.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create instruction offset (soffset) and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(int(soffset_bytes)) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer store + rocdl.RawPtrBufferStoreOp( + data, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) diff --git a/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/meta.py b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/meta.py new file mode 100644 index 000000000..23e4c849a --- /dev/null +++ b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/meta.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +import inspect +from functools import wraps + +from flydsl._mlir import ir + + +def _to_raw_value(obj): + if isinstance(obj, ir.Value): + return obj + if isinstance(obj, type): + return obj + if hasattr(obj, "__extract_to_ir_values__"): + values = obj.__extract_to_ir_values__() + if len(values) != 1: + raise ValueError(f"Primitive function expects 1 value, got {len(values)}") + return values[0] + if isinstance(obj, tuple): + return tuple(_to_raw_value(e) for e in obj) + if isinstance(obj, list): + return [_to_raw_value(e) for e in obj] + return obj + + +def _flatten_args(args, kwargs): + new_args = tuple(_to_raw_value(a) for a in args) + new_kwargs = {k: _to_raw_value(v) if k not in ("loc", "ip") else v for k, v in kwargs.items()} + return new_args, new_kwargs + + +def _caller_location(depth=1): + """Build an MLIR Location from the Python call-site *depth* frames up.""" + frame = inspect.currentframe() + for _ in range(depth + 1): + if frame is not None: + frame = frame.f_back + if frame is None: + return ir.Location.unknown() + + info = inspect.getframeinfo(frame) + pos = getattr(info, "positions", None) + line = pos.lineno if pos is not None else info.lineno + col = (pos.col_offset or 0) if pos is not None else 0 + file_loc = ir.Location.file(info.filename, line, col) + + if info.code_context: + label = " ".join(ln.strip() for ln in info.code_context) + else: + label = info.function + return ir.Location.name(label, childLoc=file_loc) + + +def traced_op(op): + @wraps(op) + def wrapper(*args, **kwargs): + loc = kwargs.pop("loc", None) + if loc is None: + loc = _caller_location(depth=1) + args, kwargs = _flatten_args(args, kwargs) + with loc: + return op(*args, **kwargs) + + return wrapper diff --git a/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/vector.py b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/vector.py new file mode 100644 index 000000000..29f8289c2 --- /dev/null +++ b/tasks/torch2flydsl/hgemm_kernel/flydsl_compat/vector.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""Vector dialect helpers and re-exports. + +The ``Vector`` class itself lives in ``typing.py`` alongside other builtin +DSL types. This module re-exports it for convenience and provides thin +wrappers around upstream ``_mlir.dialects.vector`` ops. +""" + +from __future__ import annotations + +from flydsl._mlir import ir +from flydsl._mlir.dialects import vector as _vector + +# Re-export upstream dialect for ``from flydsl.expr import vector; vector.broadcast(...)`` +from flydsl._mlir.dialects.vector import * # noqa: F401,F403,E402 +from .meta import traced_op + +# Re-export Vector and friends so ``from flydsl.expr.vector import Vector`` works +from flydsl.expr.typing import ReductionOp, Vector, empty_like, full, full_like, ones_like, zeros_like # noqa: F401 + +# ═══════════════════════════════════════════════════════════════════════ +# Dialect helper wrappers (legacy, will be deprecated) +# Prefer using Vector methods or _mlir.dialects.vector directly. +# ═══════════════════════════════════════════════════════════════════════ + + +@traced_op +def from_elements(*args, loc=None, ip=None, **kwargs): + """Construct a vector from scalar elements, auto-unwrapping ArithValue wrappers.""" + from flydsl.expr import arith as _arith_ext + + if len(args) >= 2: + args = list(args) + elems = args[1] + if isinstance(elems, (list, tuple)): + args[1] = [_arith_ext.unwrap(v) for v in elems] + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + +@traced_op +def store(value, memref, indices, *, loc=None, ip=None, **kwargs): + """Vector store wrapper that accepts ArithValue/wrappers for value/indices.""" + from flydsl.expr import arith as _arith_ext + + return _vector.store( + _arith_ext.unwrap(value), + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + **kwargs, + ) + + +# ----------------------------------------------------------------------------- +# Thin wrappers for common op classes that otherwise require `.result` access. +# ----------------------------------------------------------------------------- + + +@traced_op +def extract(vector, static_position=None, dynamic_position=None, *, loc=None, ip=None): + """Wrapper around `vector.ExtractOp(...).result`. + + When only ``dynamic_position`` is supplied (without explicit + ``static_position``), each dynamic index needs a corresponding + ``kDynamic`` sentinel in the static attribute so the ODS builder + pairs them correctly. This wrapper fills in the sentinels + automatically. + """ + from flydsl.expr import arith as _arith_ext + + if static_position is None: + static_position = [] + if dynamic_position is None: + dynamic_position = [] + dynamic_position = [_arith_ext.unwrap(i, index=True) for i in dynamic_position] + + n_static = len(static_position) + n_dynamic = len(dynamic_position) + if n_dynamic > 0 and n_static < n_dynamic: + kDynamic = ir.ShapedType.get_dynamic_size() + static_position = list(static_position) + [kDynamic] * (n_dynamic - n_static) + + return _vector.ExtractOp( + _arith_ext.unwrap(vector), + static_position=static_position, + dynamic_position=dynamic_position, + loc=loc, + ip=ip, + ).result + + +@traced_op +def load_op(result_type, memref, indices, *, loc=None, ip=None): + """Wrapper around `vector.LoadOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.LoadOp( + result_type, + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + ).result + + +@traced_op +def bitcast(result_type, source, *, loc=None, ip=None): + """Wrapper around `vector.BitCastOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.BitCastOp( + result_type, + _arith_ext.unwrap(source), + loc=loc, + ip=ip, + ).result diff --git a/tasks/torch2flydsl/hgemm_kernel/kernel.py b/tasks/torch2flydsl/hgemm_kernel/kernel.py index d4bc4b5b0..303ce1575 100644 --- a/tasks/torch2flydsl/hgemm_kernel/kernel.py +++ b/tasks/torch2flydsl/hgemm_kernel/kernel.py @@ -29,14 +29,13 @@ from flydsl.compiler.protocol import extract_to_ir_values from flydsl.expr import ( arith, - buffer_ops, const_expr, gpu, ptrtoint, range_constexpr, rocdl, - vector, ) +from flydsl_compat import buffer_ops, vector from flydsl.expr.typing import T from flydsl.runtime.device import get_rocm_arch from flydsl.utils.smem_allocator import SMEM_CAPACITY_MAP, SmemAllocator, SmemPtr diff --git a/tasks/torch2flydsl/hgemm_kernel/scripts/task_actions.py b/tasks/torch2flydsl/hgemm_kernel/scripts/task_actions.py index 79c015022..ac5923bf2 100644 --- a/tasks/torch2flydsl/hgemm_kernel/scripts/task_actions.py +++ b/tasks/torch2flydsl/hgemm_kernel/scripts/task_actions.py @@ -12,4 +12,12 @@ def check(h): raise RuntimeError("Correctness/output-contract check failed") def performance(h): - return h.arena_benchmark(warmup=10, iters=100, verbose=True) + rows = h.arena_benchmark(warmup=10, iters=100, verbose=True) + for row in rows: + # Preserve the measured values and expose the existing post-timing + # checks through the validator's event-timing evidence contract. + row["device_timing"] = { + key: value for key, value in row.items() if key.startswith("benchmark_") + } + row["timed_output_checked"] = row.get("timed_output_correctness") == "PASS" + return rows diff --git a/tasks/torch2flydsl/hgemm_kernel/task_runtime.py b/tasks/torch2flydsl/hgemm_kernel/task_runtime.py index 995831fc0..94843bf1e 100644 --- a/tasks/torch2flydsl/hgemm_kernel/task_runtime.py +++ b/tasks/torch2flydsl/hgemm_kernel/task_runtime.py @@ -117,7 +117,7 @@ def check_dependencies(paths, final_language=True): calls in the declared function; it never exposes an AITER module object. Initial Triton / provided baseline evaluation keeps its original backend. """ - forbidden = {"src", "agents", "model", "test_kernel_harness", "task_runtime", "task_reference", "task_baseline", "reference_controls", "scripts"} + forbidden = {"src", "agents", "model", "test_kernel_harness", "task_runtime", "task_reference", "task_baseline", "reference_controls", "scripts", "inspect", "gc", "builtins", "importlib"} external = {"triton", "cupy", "numba", "aiter", "ctypes", "subprocess"} loaders = {"eval", "exec", "__import__", "builtins.eval", "builtins.exec", "builtins.__import__", "importlib.import_module", "importlib.util.spec_from_file_location"} @@ -181,8 +181,57 @@ def dotted(node): if isinstance(node, ast.Attribute): return dotted(node.value) + "." + node.attr return "" + # Track simple aliases of imported objects before candidate import, so + # t = torch cannot turn a module mutation into an apparently local write. + for _ in range(len(list(ast.walk(tree)))): + changed = False + for binding in ast.walk(tree): + if not isinstance(binding, (ast.Assign, ast.AnnAssign)): + continue + value = binding.value + name = dotted(value) + if not name or name.split(".")[0] not in {v.split(".")[0] for v in aliases.values()}: + continue + targets = binding.targets if isinstance(binding, ast.Assign) else [binding.target] + for target in targets: + if isinstance(target, ast.Name) and target.id not in aliases: + aliases[target.id] = name + changed = True + if not changed: + break + imported_roots = {name.split(".")[0] for name in aliases.values()} for node in ast.walk(tree): name = dotted(node) + if name in {"sys.modules", "sys._getframe", "sys.settrace", "sys.setprofile"}: + raise ValueError(f"Protected runtime state in candidate: {name}") + if isinstance(node, ast.Attribute) and node.attr in { + "__dict__", "__globals__", "__builtins__", "__code__", "__closure__", + "__subclasses__", "__getattribute__", "f_globals", "f_locals", "f_back", + }: + raise ValueError(f"Runtime introspection is not allowed: {node.attr}") + if isinstance(node, (ast.Attribute, ast.Subscript)) and isinstance(node.ctx, (ast.Store, ast.Del)): + base = node.value + while isinstance(base, (ast.Attribute, ast.Subscript)): + base = base.value + if dotted(base).split(".")[0] in imported_roots: + raise ValueError(f"Protected dependency mutation in candidate: {name}") + if isinstance(node, ast.Name) and isinstance(node.ctx, ast.Load): + if node.id in {"globals", "locals", "vars", "setattr", "delattr"}: + raise ValueError(f"Dynamic runtime introspection is not allowed: {node.id}") + if node.id == "getattr" and not ( + isinstance(parents.get(node), ast.Call) and parents[node].func is node + ): + raise ValueError("Dynamic runtime introspection alias is not allowed: getattr") + if isinstance(node, ast.Call): + called = dotted(node.func) + if called in {"globals", "locals", "vars", "setattr", "delattr"}: + raise ValueError(f"Dynamic runtime introspection is not allowed: {called}") + if called == "getattr" and not ( + len(node.args) in (2, 3) and isinstance(node.args[1], ast.Constant) + and node.args[1].value == "_cf" + and dotted(node.args[0]).split(".")[0] not in imported_roots + ): + raise ValueError("Dynamic runtime introspection is not allowed: getattr") if final_language and isinstance(getattr(node, "ctx", None), ast.Load): if name in loaders: raise ValueError(f"Dynamic implementation loading is not allowed: {name}") diff --git a/tasks/torch2flydsl/jagged_dense_bmm_kernel/README.md b/tasks/torch2flydsl/jagged_dense_bmm_kernel/README.md index bd862b9a0..ce79d4867 100644 --- a/tasks/torch2flydsl/jagged_dense_bmm_kernel/README.md +++ b/tasks/torch2flydsl/jagged_dense_bmm_kernel/README.md @@ -75,8 +75,11 @@ and biases outside timing, poison output and replay the same measured call. Restore all inputs afterwards. The prepared function, padded output allocation, fixed metadata,10external warmups/100samples, diagnostic reference10warmups and graph policy are unchanged. No metadata construction moves into timed work. -The unchanged source uses older FlyDSL APIs; record its pinned compatible image -in full GPU qualification, and do not infer support for another runtime. +The task-local `flydsl_compat` package supplies the removed buffer helpers on +FlyDSL 0.3.2 and selects the installed helpers on older runtimes. MLIR type +construction uses the current scalar type properties. The bounded +store descriptor, original five workloads, numerical gate and timing scope are +unchanged. Runtime compatibility requires fresh GPU qualification. The candidate audit permits the original launch-metadata calculation: slices diff --git a/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/LICENSE b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/LICENSE new file mode 100644 index 000000000..c73e2f2c4 --- /dev/null +++ b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/LICENSE @@ -0,0 +1,17 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + Copyright 2025 FlyDSL Project Contributors + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/SOURCE.md b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/SOURCE.md new file mode 100644 index 000000000..83ea5ad77 --- /dev/null +++ b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/SOURCE.md @@ -0,0 +1,13 @@ +# Pinned FlyDSL compatibility helpers + +Derived from ROCm/FlyDSL commit `28a18d328b4882c999864b2df2f8f9fe3fcc8b47`, +`python/flydsl/expr/{buffer_ops,vector,meta}.py`, under Apache-2.0 (see LICENSE). +The original buffer descriptor, byte offset, masking and cache policy is retained. +Relative dependency imports now target the installed package. Removed memref +pointer extraction uses the current typed iterator and LLVM pointer conversion. +Legacy runtimes continue using their installed original helpers. No runtime +monkeypatching, external repository imports or downloads are performed. + +Current ROCDL load/store operations receive the same cache-policy bits through +their `aux` attribute instead of a removed SSA operand. Vector unwrapping uses +the current signature while retaining the surrounding MLIR location context. diff --git a/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/__init__.py b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/__init__.py new file mode 100644 index 000000000..1bbff9ffa --- /dev/null +++ b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/__init__.py @@ -0,0 +1,13 @@ +"""Task-local legacy API fallback; never modifies the installed FlyDSL package.""" +try: + import flydsl.expr.buffer_ops as buffer_ops +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.buffer_ops": + raise + from . import buffer_ops +try: + import flydsl.expr.vector as vector +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.vector": + raise + from . import vector diff --git a/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/buffer_ops.py b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/buffer_ops.py new file mode 100644 index 000000000..ab27dd821 --- /dev/null +++ b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/buffer_ops.py @@ -0,0 +1,603 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""AMD Buffer Load/Store Operations - High-level Python API + +This module provides high-level Python wrappers for AMD CDNA3/CDNA4 buffer operations. +Buffer operations use a scalar base pointer and per-thread offsets for efficient memory access. + +Example: + >>> from flydsl._mlir_helpers import buffer_ops + >>> from flydsl._mlir_helpers import arith + >>> import _mlir.extras.types as T + >>> + >>> # Create buffer resource from memref + >>> rsrc = buffer_ops.create_buffer_resource(A) + >>> + >>> # Compute offset + >>> offset = row * arith.index(4096) + col + >>> + >>> # Buffer load (4xf32) + >>> data = buffer_ops.buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Buffer store + >>> buffer_ops.buffer_store(data, rsrc, offset) +""" + +from typing import Optional, Union + +from flydsl._mlir import ir +from flydsl._mlir.dialects import arith as std_arith +from flydsl._mlir.dialects import llvm, rocdl +from flydsl._mlir.extras import types as T +from flydsl.runtime.device import is_rdna_arch +from .meta import traced_op + + +def _get_buffer_flags(arch=None): + """Get AMD buffer resource descriptor (V#) flags word (bits 127:96). + + Constructs the 32-bit flags field for rocdl.make.buffer.rsrc, following the + same logic as LLVM's AMDGPUToROCDL makeBufferRsrc(): + https://github.com/llvm/llvm-project/blob/main/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp + + Bit layout (common to all architectures): + bits [11:0] - DST_SEL: ignored by raw buffer intrinsics + bits [14:12] - DATA_FORMAT: must be nonzero, 7 = float + bits [18:15] - NUM_FORMAT: must be nonzero, 4 = 32-bit + bit [19] - In nested heap (0) + bit [20] - Behavior on unmap (0 = return 0 / ignore) + bits [22:21] - Index stride for swizzles (0) + bit [23] - Add thread ID (0) + bit [24] - Reserved: must be 1 on RDNA, 0 on CDNA + bits [26:25] - Reserved (0) + bit [27] - Non-volatile (CDNA only, 0) + bits [29:28] - OOB_SELECT (RDNA only): 0=structured, 2=none, 3=check offset + bits [31:30] - Type (must be 0) + + CDNA (gfx9xx): (7 << 12) | (4 << 15) = 0x20070 + RDNA (gfx10+): (7 << 12) | (4 << 15) | (1 << 24) | (2 << 28) = 0x21020070 + - bit 24 set to 1 (required on RDNA) + - OOB_SELECT=2 (no bounds checking, matching LLVM boundsCheck=false) + """ + import os + + if arch is None: + arch = os.environ.get("FLYDSL_GPU_ARCH") + flags = (7 << 12) | (4 << 15) + if is_rdna_arch(arch): + flags |= 1 << 24 # reserved bit, must be 1 on RDNA + flags |= 2 << 28 # OOB_SELECT = 2 (no bounds checking) + return flags + + +__all__ = [ + "create_llvm_ptr", + "get_element_ptr", + "create_buffer_resource", + "create_buffer_resource_from_addr", + "buffer_load", + "buffer_store", + "BufferResourceDescriptor", + "extract_base_index", +] + + +def _unwrap_value(value): + """Recursively unwrap ArithValue or similar wrappers to get the actual MLIR value. + + Handles: + - FlyDSL ArithValue (has ._value) + - flyc DSL Numeric like fx.Int32 (has .ir_value() method) + - flyc ArithValue (is already ir.Value subclass) + """ + # DSL Numeric (Int32, Float32, etc.) — use ir_value() to materialize + if hasattr(value, "ir_value") and not isinstance(value, ir.Value): + return value.ir_value() + max_depth = 10 # Safety limit + depth = 0 + while depth < max_depth and not isinstance(value, ir.Value): + if hasattr(value, "_value"): + value = value._value + elif hasattr(value, "value"): + value = value.value + else: + break + depth += 1 + return value + + +def _create_i32_constant(value: int) -> ir.Value: + """Create i32 constant using standard MLIR arith dialect.""" + i32_type = T.i32() + if value > 0x7FFFFFFF: + value = int(value - 2**32) + attr = ir.IntegerAttr.get(i32_type, value) + op = std_arith.ConstantOp(i32_type, attr) + return _unwrap_value(op.result) + + +def _create_i16_constant(value: int) -> ir.Value: + """Create i16 constant using standard MLIR arith dialect.""" + i16_type = T.i16() + attr = ir.IntegerAttr.get(i16_type, value) + op = std_arith.ConstantOp(i16_type, attr) + return _unwrap_value(op.result) + + +def _create_i64_constant(value: int) -> ir.Value: + """Create i64 constant using standard MLIR arith dialect.""" + i64_type = T.i64() + attr = ir.IntegerAttr.get(i64_type, value) + op = std_arith.ConstantOp(i64_type, attr) + return _unwrap_value(op.result) + + +def create_llvm_ptr(value, address_space: int = 0) -> ir.Value: + """Create an LLVM pointer from an integer or index value.""" + value = _unwrap_value(value) + if isinstance(value.type, ir.IndexType): + i64_type = T.i64() + value = _unwrap_value(std_arith.IndexCastOp(i64_type, value).result) + ptr_type = ir.Type.parse(f"!llvm.ptr<{address_space}>") + return llvm.IntToPtrOp(ptr_type, value).result + + +def extract_base_index(tensor, address_space: int = 1) -> ir.Value: + """Extract the base address of a fly.memref as an index value. + + Inverse of :func:`create_llvm_ptr` (index -> ptr). Useful when ISA + requires a raw pointer instead of a buffer resource descriptor + (e.g. global_atomic_pk_add_bf16 on gfx942). + """ + from flydsl._mlir.dialects import fly as _fly + from flydsl._mlir.dialects import memref as _memref + + raw = _unwrap_value(tensor) + try: + ir.MemRefType(raw.type) + return _memref.extract_aligned_pointer_as_index(raw) + except ValueError: + pass + + # FlyDSL 0.3 uses a typed iterator instead of the removed extract op. + from flydsl.expr import get_iter, to_llvm_ptr + ptr = to_llvm_ptr(get_iter(raw)) + i64_val = llvm.PtrToIntOp(ir.IntegerType.get_signless(64), ptr).result + return _unwrap_value(std_arith.IndexCastOp(ir.IndexType.get(), i64_val).result) + + +def get_element_ptr( + base_ptr, + byte_offset: Union[int, ir.Value, None] = None, + static_byte_offset: int = 0, + elem_type: Optional[ir.Type] = None, + no_wrap_flags=None, +) -> ir.Value: + """Build an LLVM GEP from a base pointer plus byte offsets.""" + _gep_dynamic_index_sentinel = -(2**31) + + base_ptr = _unwrap_value(base_ptr) + if not isinstance(static_byte_offset, int): + raise TypeError(f"static_byte_offset must be int, got {type(static_byte_offset).__name__}") + if elem_type is None: + elem_type = T.i8() + elif callable(elem_type): + elem_type = elem_type() + + if byte_offset is None: + dynamic_indices = [] + raw_constant_indices = [int(static_byte_offset)] + elif isinstance(byte_offset, int): + dynamic_indices = [] + raw_constant_indices = [int(byte_offset) + int(static_byte_offset)] + else: + offset_val = _unwrap_value(byte_offset) + if isinstance(offset_val.type, ir.IndexType): + i64_type = T.i64() + offset_val = _unwrap_value(std_arith.IndexCastOp(i64_type, offset_val).result) + elif not isinstance(offset_val.type, ir.IntegerType): + raise TypeError("byte_offset must be int, index, or integer-typed MLIR value; " f"got {offset_val.type}") + + if static_byte_offset != 0: + static_type = offset_val.type + static_attr = ir.IntegerAttr.get(static_type, int(static_byte_offset)) + static_const = _unwrap_value(std_arith.ConstantOp(static_type, static_attr).result) + offset_val = _unwrap_value(std_arith.AddIOp(offset_val, static_const).result) + + dynamic_indices = [offset_val] + raw_constant_indices = [_gep_dynamic_index_sentinel] + + return llvm.GEPOp( + base_ptr.type, + base_ptr, + dynamic_indices, + raw_constant_indices, + elem_type, + no_wrap_flags, + ).result + + +class BufferResourceDescriptor: + """AMD Buffer Resource Descriptor + + A buffer resource descriptor contains: + - base_pointer: Scalar base pointer (wave-uniform, stored in SGPRs) + - stride: Stride for structured buffers (typically 0 for contiguous) + - num_records: Buffer size in bytes + - flags: Data format and access flags + + The descriptor is stored in a special LLVM pointer type (!llvm.ptr<8>) + """ + + def __init__(self, rsrc: ir.Value): + """Initialize with ROCDL resource descriptor value.""" + self.rsrc = rsrc + + @staticmethod + def from_memref( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + data_format: str = "f32", + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, + ) -> "BufferResourceDescriptor": + """Create buffer resource descriptor from memref. + + Args: + memref_val: Memref value to create descriptor for + stride: Stride in elements (0 for contiguous) + max_size: If True, use max buffer size for flexibility + num_records_bytes: Override buffer size (in BYTES) used by hardware OOB checking. + If provided, this takes precedence over `max_size`. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + data_format: Data format ('f32', 'f16', 'i32', etc.) + + Returns: + BufferResourceDescriptor instance + + Example: + >>> rsrc = BufferResourceDescriptor.from_memref(A) + """ + # Extract raw pointer from fly.memref. + raw_val = _unwrap_value(memref_val) + from flydsl._mlir.dialects import fly as _fly + + # Preserve byte offsets and descriptor policy with the current pointer API. + from flydsl.expr import get_iter, to_llvm_ptr + base_ptr = to_llvm_ptr(get_iter(raw_val)) + if base_byte_offset is not None: + base_ptr = get_element_ptr(base_ptr, byte_offset=base_byte_offset) + + # Create buffer resource descriptor + flags_val = _get_buffer_flags() + flags = _create_i32_constant(flags_val) + stride_val = _create_i16_constant(stride) + + def _num_records_from_memref_type() -> Optional[int]: + """Best-effort: derive logical buffer size (in bytes) from static memref type.""" + try: + mt = ir.MemRefType(_unwrap_value(memref_val).type) + shape = list(mt.shape) + if any(int(d) < 0 for d in shape): + return None + # Compute element size in bytes (scalar element type). + elem_t = mt.element_type + elem_bits = getattr(elem_t, "width", None) + if elem_bits is None: + return None + elem_bytes = int(elem_bits) // 8 + if elem_bytes <= 0: + return None + num_elems = 1 + for d in shape: + num_elems *= int(d) + return int(num_elems) * int(elem_bytes) + except Exception: + return None + + if num_records_bytes is not None: + # Caller-provided size in BYTES (preferred for exact hardware OOB behavior). + if isinstance(num_records_bytes, int): + nbytes = int(num_records_bytes) + if nbytes <= 0: + nbytes = 0 + # Descriptor uses i32 bytes; clamp to the max representable. + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(nbytes) + else: + v = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(v.type, ir.IntegerType) or v.type.width != 64: + if isinstance(v.type, ir.IndexType): + op = std_arith.IndexCastOp(i64_type, v) + else: + op = std_arith.ExtSIOp(i64_type, v) + v = _unwrap_value(op.result) + num_records = v + elif max_size: + # Use max for flexibility (hardware will check actual bounds) + # Note: FlyDSL's rocdl.make.buffer.rsrc requires i32, not i64 + num_records = _create_i64_constant(0xFFFFFFFF) # FALLBACK_MAX_SIZE + else: + # Use the logical memref size (in bytes) for hardware OOB checking. + nbytes = _num_records_from_memref_type() + if nbytes is None: + # Fall back to max-size if we can't infer statically. + num_records = _create_i64_constant(0xFFFFFFFF) + else: + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(int(nbytes)) + + # Create resource descriptor (returns !llvm.ptr<8>) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + rsrc = rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride_val, num_records, flags).result + + return BufferResourceDescriptor(rsrc) + + +def create_buffer_resource_from_addr( + addr_i64: ir.Value, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from a raw i64 device address. + + Useful when working with runtime pointer arrays (e.g. IPC-mapped addresses + or device-side pointer tables) where no fly.memref is available. + The full address is encoded as the buffer base; callers should pass + byte offset 0 to buffer_load / buffer_store. + + Args: + addr_i64: Raw 64-bit device address (i64 MLIR value). + num_records_bytes: Optional buffer size in bytes for hardware OOB checking. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>). + + Example: + >>> rsrc = create_buffer_resource_from_addr(raw_addr_i64) + >>> data = buffer_load(rsrc, i32_zero, vec_width=4, dtype=T.i32) + """ + addr_i64 = _unwrap_value(addr_i64) + ptr_type = ir.Type.parse("!llvm.ptr") + base_ptr = llvm.IntToPtrOp(ptr_type, addr_i64).result + flags = _create_i32_constant(_get_buffer_flags()) + stride = _create_i16_constant(0) + if num_records_bytes is None: + num_records = _create_i64_constant(0xFFFFFFFF) + elif isinstance(num_records_bytes, int): + nbytes = max(0, min(int(num_records_bytes), 0xFFFFFFFF)) + num_records = _create_i64_constant(nbytes) + else: + num_records = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(num_records.type, ir.IntegerType) or num_records.type.width != 64: + if isinstance(num_records.type, ir.IndexType): + num_records = _unwrap_value(std_arith.IndexCastOp(i64_type, num_records).result) + else: + num_records = _unwrap_value(std_arith.ExtSIOp(i64_type, num_records).result) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + return rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride, num_records, flags).result + + +@traced_op +def create_buffer_resource( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from memref. + + This is a simplified wrapper around BufferResourceDescriptor.from_memref() + that returns the raw ROCDL resource value. + + Args: + memref_val: Memref value + stride: Buffer stride (0 for contiguous) + max_size: Use maximum buffer size + num_records_bytes: Override buffer size in bytes. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>) + + Example: + >>> rsrc = create_buffer_resource(A) + >>> data = buffer_load(rsrc, offset) + """ + desc = BufferResourceDescriptor.from_memref( + memref_val, + stride, + max_size, + num_records_bytes=num_records_bytes, + base_byte_offset=base_byte_offset, + ) + return desc.rsrc + + +@traced_op +def buffer_load( + rsrc: ir.Value, + offset: ir.Value, + vec_width: int = 4, + dtype=None, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + soffset_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """AMD buffer load operation. + + Load data from global memory using buffer descriptor and offset. + Uses hardware-level bounds checking and vectorization. + + Args: + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + vec_width: Vector width (1, 2, or 4) + dtype: Element data type (None for f32, or ir.F32Type, etc.) + mask: Optional mask for predicated load (i1 type) + cache_modifier: Cache control flags (0 for default) + soffset_bytes: Optional scalar offset (in BYTES) added by the buffer instruction (soffset). + Use this to fold small constant deltas into the instruction instead of emitting + extra VGPR address arithmetic. + + Returns: + Loaded data (scalar or vector depending on vec_width) + + Example: + >>> # Load 4xf32 + >>> data = buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Load with mask + >>> data = buffer_load(rsrc, offset, vec_width=4, mask=valid) + """ + # Default dtype to f32 + if dtype is None: + dtype = T.f32() + # Accept DSL Numeric class (e.g. fx.Int32) as dtype: unwrap to ir.Type + elif hasattr(dtype, "ir_type"): + dtype = dtype.ir_type + + # Unwrap offset first (accept Python ints and DSL Numeric values). + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: Buffer load offset is in BYTES, not elements! + # For vec4xf32, each element is 4 bytes, so multiply offset by 4 + element_bytes = dtype.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create vector type + if vec_width == 1: + result_type = dtype + else: + result_type = ir.VectorType.get([vec_width], dtype) + + # Create instruction offset and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(soffset_bytes) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer load + load_op = rocdl.RawPtrBufferLoadOp( + result_type, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) + + return load_op.result + + +@traced_op +def buffer_store( + data: ir.Value, + rsrc: ir.Value, + offset: ir.Value, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + *, + soffset_bytes: Optional[Union[int, ir.Value]] = None, + offset_is_bytes: bool = False, +): + """AMD buffer store operation. + + Store data to global memory using buffer descriptor and offset. + + Args: + data: Data to store (scalar or vector) + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + mask: Optional mask for predicated store (i1 type) + cache_modifier: Cache control flags (0 for default) + + Example: + >>> buffer_store(data, rsrc, offset) + >>> + >>> # Store with mask + >>> buffer_store(data, rsrc, offset, mask=valid) + """ + # Unwrap all inputs (accept DSL Numeric values via ir_value()) + if hasattr(data, "ir_value"): + data = data.ir_value() + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + data = _unwrap_value(data) + rsrc = _unwrap_value(rsrc) + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: RawPtrBufferStoreOp offset is in BYTES. + # For backward compat, `buffer_store()` accepts element offsets by default + # and scales them to bytes. Set `offset_is_bytes=True` to skip scaling. + if not offset_is_bytes: + # Get element size from data type + data_type = data.type + if hasattr(data_type, "element_type"): # Vector type + element_type = data_type.element_type + else: # Scalar type + element_type = data_type + element_bytes = element_type.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create instruction offset (soffset) and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(int(soffset_bytes)) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer store + rocdl.RawPtrBufferStoreOp( + data, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) diff --git a/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/meta.py b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/meta.py new file mode 100644 index 000000000..23e4c849a --- /dev/null +++ b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/meta.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +import inspect +from functools import wraps + +from flydsl._mlir import ir + + +def _to_raw_value(obj): + if isinstance(obj, ir.Value): + return obj + if isinstance(obj, type): + return obj + if hasattr(obj, "__extract_to_ir_values__"): + values = obj.__extract_to_ir_values__() + if len(values) != 1: + raise ValueError(f"Primitive function expects 1 value, got {len(values)}") + return values[0] + if isinstance(obj, tuple): + return tuple(_to_raw_value(e) for e in obj) + if isinstance(obj, list): + return [_to_raw_value(e) for e in obj] + return obj + + +def _flatten_args(args, kwargs): + new_args = tuple(_to_raw_value(a) for a in args) + new_kwargs = {k: _to_raw_value(v) if k not in ("loc", "ip") else v for k, v in kwargs.items()} + return new_args, new_kwargs + + +def _caller_location(depth=1): + """Build an MLIR Location from the Python call-site *depth* frames up.""" + frame = inspect.currentframe() + for _ in range(depth + 1): + if frame is not None: + frame = frame.f_back + if frame is None: + return ir.Location.unknown() + + info = inspect.getframeinfo(frame) + pos = getattr(info, "positions", None) + line = pos.lineno if pos is not None else info.lineno + col = (pos.col_offset or 0) if pos is not None else 0 + file_loc = ir.Location.file(info.filename, line, col) + + if info.code_context: + label = " ".join(ln.strip() for ln in info.code_context) + else: + label = info.function + return ir.Location.name(label, childLoc=file_loc) + + +def traced_op(op): + @wraps(op) + def wrapper(*args, **kwargs): + loc = kwargs.pop("loc", None) + if loc is None: + loc = _caller_location(depth=1) + args, kwargs = _flatten_args(args, kwargs) + with loc: + return op(*args, **kwargs) + + return wrapper diff --git a/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/vector.py b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/vector.py new file mode 100644 index 000000000..29f8289c2 --- /dev/null +++ b/tasks/torch2flydsl/jagged_dense_bmm_kernel/flydsl_compat/vector.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""Vector dialect helpers and re-exports. + +The ``Vector`` class itself lives in ``typing.py`` alongside other builtin +DSL types. This module re-exports it for convenience and provides thin +wrappers around upstream ``_mlir.dialects.vector`` ops. +""" + +from __future__ import annotations + +from flydsl._mlir import ir +from flydsl._mlir.dialects import vector as _vector + +# Re-export upstream dialect for ``from flydsl.expr import vector; vector.broadcast(...)`` +from flydsl._mlir.dialects.vector import * # noqa: F401,F403,E402 +from .meta import traced_op + +# Re-export Vector and friends so ``from flydsl.expr.vector import Vector`` works +from flydsl.expr.typing import ReductionOp, Vector, empty_like, full, full_like, ones_like, zeros_like # noqa: F401 + +# ═══════════════════════════════════════════════════════════════════════ +# Dialect helper wrappers (legacy, will be deprecated) +# Prefer using Vector methods or _mlir.dialects.vector directly. +# ═══════════════════════════════════════════════════════════════════════ + + +@traced_op +def from_elements(*args, loc=None, ip=None, **kwargs): + """Construct a vector from scalar elements, auto-unwrapping ArithValue wrappers.""" + from flydsl.expr import arith as _arith_ext + + if len(args) >= 2: + args = list(args) + elems = args[1] + if isinstance(elems, (list, tuple)): + args[1] = [_arith_ext.unwrap(v) for v in elems] + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + +@traced_op +def store(value, memref, indices, *, loc=None, ip=None, **kwargs): + """Vector store wrapper that accepts ArithValue/wrappers for value/indices.""" + from flydsl.expr import arith as _arith_ext + + return _vector.store( + _arith_ext.unwrap(value), + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + **kwargs, + ) + + +# ----------------------------------------------------------------------------- +# Thin wrappers for common op classes that otherwise require `.result` access. +# ----------------------------------------------------------------------------- + + +@traced_op +def extract(vector, static_position=None, dynamic_position=None, *, loc=None, ip=None): + """Wrapper around `vector.ExtractOp(...).result`. + + When only ``dynamic_position`` is supplied (without explicit + ``static_position``), each dynamic index needs a corresponding + ``kDynamic`` sentinel in the static attribute so the ODS builder + pairs them correctly. This wrapper fills in the sentinels + automatically. + """ + from flydsl.expr import arith as _arith_ext + + if static_position is None: + static_position = [] + if dynamic_position is None: + dynamic_position = [] + dynamic_position = [_arith_ext.unwrap(i, index=True) for i in dynamic_position] + + n_static = len(static_position) + n_dynamic = len(dynamic_position) + if n_dynamic > 0 and n_static < n_dynamic: + kDynamic = ir.ShapedType.get_dynamic_size() + static_position = list(static_position) + [kDynamic] * (n_dynamic - n_static) + + return _vector.ExtractOp( + _arith_ext.unwrap(vector), + static_position=static_position, + dynamic_position=dynamic_position, + loc=loc, + ip=ip, + ).result + + +@traced_op +def load_op(result_type, memref, indices, *, loc=None, ip=None): + """Wrapper around `vector.LoadOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.LoadOp( + result_type, + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + ).result + + +@traced_op +def bitcast(result_type, source, *, loc=None, ip=None): + """Wrapper around `vector.BitCastOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.BitCastOp( + result_type, + _arith_ext.unwrap(source), + loc=loc, + ip=ip, + ).result diff --git a/tasks/torch2flydsl/jagged_dense_bmm_kernel/kernel.py b/tasks/torch2flydsl/jagged_dense_bmm_kernel/kernel.py index 5edfd9902..275a53a1a 100644 --- a/tasks/torch2flydsl/jagged_dense_bmm_kernel/kernel.py +++ b/tasks/torch2flydsl/jagged_dense_bmm_kernel/kernel.py @@ -29,6 +29,8 @@ import flydsl.compiler as flyc import flydsl.expr as fx +from flydsl_compat import buffer_ops +from flydsl._mlir import ir BLOCK_M = 128 BLOCK_N = 128 @@ -49,7 +51,6 @@ def make_bounded_buffer_tensor(tensor, num_records_bytes): hardware OOB-drops stores past num_records_bytes. Mirrors the installed make_buffer_tensor body (which hardcodes max_size).""" from flydsl._mlir.dialects.fly_rocdl import TargetAddressSpace - from flydsl.expr.buffer_ops import _get_buffer_flags elem_ty = tensor.element_type ptr = fx.get_iter(tensor) @@ -65,7 +66,7 @@ def make_bounded_buffer_tensor(tensor, num_records_bytes): ptr, fx.Int16(0).ir_value(), num_records_bytes.ir_value(), - fx.Int32(_get_buffer_flags()).ir_value(), + fx.Int32(buffer_ops._get_buffer_flags()).ir_value(), ], ) return fx.make_view(buf_ptr, layout) @@ -92,20 +93,20 @@ def jdbba_kernel( block_n_idx = pid_mn % N_BLOCKS # --- Device group resolution (read seq_offsets[b], seq_offsets[b+1]) --- - seq_rsrc = fx.buffer_ops.create_buffer_resource(SEQ_OFFSETS, max_size=True) - seq_start = fx.buffer_ops.buffer_load( - seq_rsrc, fx.Int32(off_b), vec_width=1, dtype=fx.T.i32() + seq_rsrc = buffer_ops.create_buffer_resource(SEQ_OFFSETS, max_size=True) + seq_start = buffer_ops.buffer_load( + seq_rsrc, fx.Int32(off_b), vec_width=1, dtype=fx.Int32.ir_type ) - seq_end = fx.buffer_ops.buffer_load( - seq_rsrc, fx.Int32(off_b) + fx.Int32(1), vec_width=1, dtype=fx.T.i32() + seq_end = buffer_ops.buffer_load( + seq_rsrc, fx.Int32(off_b) + fx.Int32(1), vec_width=1, dtype=fx.Int32.ir_type ) # seq_start/seq_end are block-uniform (one group per block) but buffer_load # types them per-lane (VGPR). Scalarize so everything derived from them -- # M_b, the A/C/B/bias base offsets, and the C buffer-descriptor bound -- is # uniform (SGPR). Otherwise the divergent C descriptor forces the epilogue # store into a per-lane readfirstlane/exec-mask waterfall. - seq_start = fx.rocdl.readfirstlane(fx.T.i32(), seq_start) - seq_end = fx.rocdl.readfirstlane(fx.T.i32(), seq_end) + seq_start = fx.rocdl.readfirstlane(fx.Int32.ir_type, seq_start) + seq_end = fx.rocdl.readfirstlane(fx.Int32.ir_type, seq_end) M_b = seq_end - seq_start start_m = fx.Int32(block_m_idx) * fx.Int32(BLOCK_M) @@ -238,10 +239,10 @@ def run_pipeline_stage(read_stage, next_k, read_next=True): thr_gBias = thr_copy_r2g_C.partition_S(gBias) bias_frag = fx.make_fragment_like(thr_gBias) fx.copy(fx.make_copy_atom(fx.rocdl.BufferCopy16b(), fx.BFloat16), thr_gBias, bias_frag) - bias_f32 = fx.arith.ExtFOp(fx.T.VectorType.get([64], fx.T.f32()), bias_frag.load()).result + bias_f32 = fx.arith.ExtFOp(ir.VectorType.get([64], fx.Float32.ir_type), bias_frag.load()).result mma_frag_C_bf16.store( fx.arith.trunc_f( - fx.T.VectorType.get([64], fx.T.bf16()), + ir.VectorType.get([64], fx.BFloat16.ir_type), fx.arith.addf(mma_frag_C.load(), bias_f32), ) ) diff --git a/tasks/torch2flydsl/moe_sorting_kernel/README.md b/tasks/torch2flydsl/moe_sorting_kernel/README.md index eb8b9e883..7347fe003 100644 --- a/tasks/torch2flydsl/moe_sorting_kernel/README.md +++ b/tasks/torch2flydsl/moe_sorting_kernel/README.md @@ -75,3 +75,8 @@ plan still corresponds to the original case. Preserve the original10external warmups/100samples and graph policy. Candidate operator calls are audited outside timing for FlyDSL execution and permitted host preparation. The unchanged source uses older FlyDSL APIs; report the pinned compatible runtime used for validation. + +The task-local `flydsl_compat` helpers preserve the legacy buffer/vector API +when the installed FlyDSL no longer supplies it. See `flydsl_compat/SOURCE.md` +for the pinned upstream source and retained license. Workloads, numerical gates, +and timing parameters are unchanged. diff --git a/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/LICENSE b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/LICENSE new file mode 100644 index 000000000..c73e2f2c4 --- /dev/null +++ b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/LICENSE @@ -0,0 +1,17 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + Copyright 2025 FlyDSL Project Contributors + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/SOURCE.md b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/SOURCE.md new file mode 100644 index 000000000..83ea5ad77 --- /dev/null +++ b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/SOURCE.md @@ -0,0 +1,13 @@ +# Pinned FlyDSL compatibility helpers + +Derived from ROCm/FlyDSL commit `28a18d328b4882c999864b2df2f8f9fe3fcc8b47`, +`python/flydsl/expr/{buffer_ops,vector,meta}.py`, under Apache-2.0 (see LICENSE). +The original buffer descriptor, byte offset, masking and cache policy is retained. +Relative dependency imports now target the installed package. Removed memref +pointer extraction uses the current typed iterator and LLVM pointer conversion. +Legacy runtimes continue using their installed original helpers. No runtime +monkeypatching, external repository imports or downloads are performed. + +Current ROCDL load/store operations receive the same cache-policy bits through +their `aux` attribute instead of a removed SSA operand. Vector unwrapping uses +the current signature while retaining the surrounding MLIR location context. diff --git a/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/__init__.py b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/__init__.py new file mode 100644 index 000000000..1bbff9ffa --- /dev/null +++ b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/__init__.py @@ -0,0 +1,13 @@ +"""Task-local legacy API fallback; never modifies the installed FlyDSL package.""" +try: + import flydsl.expr.buffer_ops as buffer_ops +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.buffer_ops": + raise + from . import buffer_ops +try: + import flydsl.expr.vector as vector +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.vector": + raise + from . import vector diff --git a/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/buffer_ops.py b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/buffer_ops.py new file mode 100644 index 000000000..ab27dd821 --- /dev/null +++ b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/buffer_ops.py @@ -0,0 +1,603 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""AMD Buffer Load/Store Operations - High-level Python API + +This module provides high-level Python wrappers for AMD CDNA3/CDNA4 buffer operations. +Buffer operations use a scalar base pointer and per-thread offsets for efficient memory access. + +Example: + >>> from flydsl._mlir_helpers import buffer_ops + >>> from flydsl._mlir_helpers import arith + >>> import _mlir.extras.types as T + >>> + >>> # Create buffer resource from memref + >>> rsrc = buffer_ops.create_buffer_resource(A) + >>> + >>> # Compute offset + >>> offset = row * arith.index(4096) + col + >>> + >>> # Buffer load (4xf32) + >>> data = buffer_ops.buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Buffer store + >>> buffer_ops.buffer_store(data, rsrc, offset) +""" + +from typing import Optional, Union + +from flydsl._mlir import ir +from flydsl._mlir.dialects import arith as std_arith +from flydsl._mlir.dialects import llvm, rocdl +from flydsl._mlir.extras import types as T +from flydsl.runtime.device import is_rdna_arch +from .meta import traced_op + + +def _get_buffer_flags(arch=None): + """Get AMD buffer resource descriptor (V#) flags word (bits 127:96). + + Constructs the 32-bit flags field for rocdl.make.buffer.rsrc, following the + same logic as LLVM's AMDGPUToROCDL makeBufferRsrc(): + https://github.com/llvm/llvm-project/blob/main/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp + + Bit layout (common to all architectures): + bits [11:0] - DST_SEL: ignored by raw buffer intrinsics + bits [14:12] - DATA_FORMAT: must be nonzero, 7 = float + bits [18:15] - NUM_FORMAT: must be nonzero, 4 = 32-bit + bit [19] - In nested heap (0) + bit [20] - Behavior on unmap (0 = return 0 / ignore) + bits [22:21] - Index stride for swizzles (0) + bit [23] - Add thread ID (0) + bit [24] - Reserved: must be 1 on RDNA, 0 on CDNA + bits [26:25] - Reserved (0) + bit [27] - Non-volatile (CDNA only, 0) + bits [29:28] - OOB_SELECT (RDNA only): 0=structured, 2=none, 3=check offset + bits [31:30] - Type (must be 0) + + CDNA (gfx9xx): (7 << 12) | (4 << 15) = 0x20070 + RDNA (gfx10+): (7 << 12) | (4 << 15) | (1 << 24) | (2 << 28) = 0x21020070 + - bit 24 set to 1 (required on RDNA) + - OOB_SELECT=2 (no bounds checking, matching LLVM boundsCheck=false) + """ + import os + + if arch is None: + arch = os.environ.get("FLYDSL_GPU_ARCH") + flags = (7 << 12) | (4 << 15) + if is_rdna_arch(arch): + flags |= 1 << 24 # reserved bit, must be 1 on RDNA + flags |= 2 << 28 # OOB_SELECT = 2 (no bounds checking) + return flags + + +__all__ = [ + "create_llvm_ptr", + "get_element_ptr", + "create_buffer_resource", + "create_buffer_resource_from_addr", + "buffer_load", + "buffer_store", + "BufferResourceDescriptor", + "extract_base_index", +] + + +def _unwrap_value(value): + """Recursively unwrap ArithValue or similar wrappers to get the actual MLIR value. + + Handles: + - FlyDSL ArithValue (has ._value) + - flyc DSL Numeric like fx.Int32 (has .ir_value() method) + - flyc ArithValue (is already ir.Value subclass) + """ + # DSL Numeric (Int32, Float32, etc.) — use ir_value() to materialize + if hasattr(value, "ir_value") and not isinstance(value, ir.Value): + return value.ir_value() + max_depth = 10 # Safety limit + depth = 0 + while depth < max_depth and not isinstance(value, ir.Value): + if hasattr(value, "_value"): + value = value._value + elif hasattr(value, "value"): + value = value.value + else: + break + depth += 1 + return value + + +def _create_i32_constant(value: int) -> ir.Value: + """Create i32 constant using standard MLIR arith dialect.""" + i32_type = T.i32() + if value > 0x7FFFFFFF: + value = int(value - 2**32) + attr = ir.IntegerAttr.get(i32_type, value) + op = std_arith.ConstantOp(i32_type, attr) + return _unwrap_value(op.result) + + +def _create_i16_constant(value: int) -> ir.Value: + """Create i16 constant using standard MLIR arith dialect.""" + i16_type = T.i16() + attr = ir.IntegerAttr.get(i16_type, value) + op = std_arith.ConstantOp(i16_type, attr) + return _unwrap_value(op.result) + + +def _create_i64_constant(value: int) -> ir.Value: + """Create i64 constant using standard MLIR arith dialect.""" + i64_type = T.i64() + attr = ir.IntegerAttr.get(i64_type, value) + op = std_arith.ConstantOp(i64_type, attr) + return _unwrap_value(op.result) + + +def create_llvm_ptr(value, address_space: int = 0) -> ir.Value: + """Create an LLVM pointer from an integer or index value.""" + value = _unwrap_value(value) + if isinstance(value.type, ir.IndexType): + i64_type = T.i64() + value = _unwrap_value(std_arith.IndexCastOp(i64_type, value).result) + ptr_type = ir.Type.parse(f"!llvm.ptr<{address_space}>") + return llvm.IntToPtrOp(ptr_type, value).result + + +def extract_base_index(tensor, address_space: int = 1) -> ir.Value: + """Extract the base address of a fly.memref as an index value. + + Inverse of :func:`create_llvm_ptr` (index -> ptr). Useful when ISA + requires a raw pointer instead of a buffer resource descriptor + (e.g. global_atomic_pk_add_bf16 on gfx942). + """ + from flydsl._mlir.dialects import fly as _fly + from flydsl._mlir.dialects import memref as _memref + + raw = _unwrap_value(tensor) + try: + ir.MemRefType(raw.type) + return _memref.extract_aligned_pointer_as_index(raw) + except ValueError: + pass + + # FlyDSL 0.3 uses a typed iterator instead of the removed extract op. + from flydsl.expr import get_iter, to_llvm_ptr + ptr = to_llvm_ptr(get_iter(raw)) + i64_val = llvm.PtrToIntOp(ir.IntegerType.get_signless(64), ptr).result + return _unwrap_value(std_arith.IndexCastOp(ir.IndexType.get(), i64_val).result) + + +def get_element_ptr( + base_ptr, + byte_offset: Union[int, ir.Value, None] = None, + static_byte_offset: int = 0, + elem_type: Optional[ir.Type] = None, + no_wrap_flags=None, +) -> ir.Value: + """Build an LLVM GEP from a base pointer plus byte offsets.""" + _gep_dynamic_index_sentinel = -(2**31) + + base_ptr = _unwrap_value(base_ptr) + if not isinstance(static_byte_offset, int): + raise TypeError(f"static_byte_offset must be int, got {type(static_byte_offset).__name__}") + if elem_type is None: + elem_type = T.i8() + elif callable(elem_type): + elem_type = elem_type() + + if byte_offset is None: + dynamic_indices = [] + raw_constant_indices = [int(static_byte_offset)] + elif isinstance(byte_offset, int): + dynamic_indices = [] + raw_constant_indices = [int(byte_offset) + int(static_byte_offset)] + else: + offset_val = _unwrap_value(byte_offset) + if isinstance(offset_val.type, ir.IndexType): + i64_type = T.i64() + offset_val = _unwrap_value(std_arith.IndexCastOp(i64_type, offset_val).result) + elif not isinstance(offset_val.type, ir.IntegerType): + raise TypeError("byte_offset must be int, index, or integer-typed MLIR value; " f"got {offset_val.type}") + + if static_byte_offset != 0: + static_type = offset_val.type + static_attr = ir.IntegerAttr.get(static_type, int(static_byte_offset)) + static_const = _unwrap_value(std_arith.ConstantOp(static_type, static_attr).result) + offset_val = _unwrap_value(std_arith.AddIOp(offset_val, static_const).result) + + dynamic_indices = [offset_val] + raw_constant_indices = [_gep_dynamic_index_sentinel] + + return llvm.GEPOp( + base_ptr.type, + base_ptr, + dynamic_indices, + raw_constant_indices, + elem_type, + no_wrap_flags, + ).result + + +class BufferResourceDescriptor: + """AMD Buffer Resource Descriptor + + A buffer resource descriptor contains: + - base_pointer: Scalar base pointer (wave-uniform, stored in SGPRs) + - stride: Stride for structured buffers (typically 0 for contiguous) + - num_records: Buffer size in bytes + - flags: Data format and access flags + + The descriptor is stored in a special LLVM pointer type (!llvm.ptr<8>) + """ + + def __init__(self, rsrc: ir.Value): + """Initialize with ROCDL resource descriptor value.""" + self.rsrc = rsrc + + @staticmethod + def from_memref( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + data_format: str = "f32", + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, + ) -> "BufferResourceDescriptor": + """Create buffer resource descriptor from memref. + + Args: + memref_val: Memref value to create descriptor for + stride: Stride in elements (0 for contiguous) + max_size: If True, use max buffer size for flexibility + num_records_bytes: Override buffer size (in BYTES) used by hardware OOB checking. + If provided, this takes precedence over `max_size`. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + data_format: Data format ('f32', 'f16', 'i32', etc.) + + Returns: + BufferResourceDescriptor instance + + Example: + >>> rsrc = BufferResourceDescriptor.from_memref(A) + """ + # Extract raw pointer from fly.memref. + raw_val = _unwrap_value(memref_val) + from flydsl._mlir.dialects import fly as _fly + + # Preserve byte offsets and descriptor policy with the current pointer API. + from flydsl.expr import get_iter, to_llvm_ptr + base_ptr = to_llvm_ptr(get_iter(raw_val)) + if base_byte_offset is not None: + base_ptr = get_element_ptr(base_ptr, byte_offset=base_byte_offset) + + # Create buffer resource descriptor + flags_val = _get_buffer_flags() + flags = _create_i32_constant(flags_val) + stride_val = _create_i16_constant(stride) + + def _num_records_from_memref_type() -> Optional[int]: + """Best-effort: derive logical buffer size (in bytes) from static memref type.""" + try: + mt = ir.MemRefType(_unwrap_value(memref_val).type) + shape = list(mt.shape) + if any(int(d) < 0 for d in shape): + return None + # Compute element size in bytes (scalar element type). + elem_t = mt.element_type + elem_bits = getattr(elem_t, "width", None) + if elem_bits is None: + return None + elem_bytes = int(elem_bits) // 8 + if elem_bytes <= 0: + return None + num_elems = 1 + for d in shape: + num_elems *= int(d) + return int(num_elems) * int(elem_bytes) + except Exception: + return None + + if num_records_bytes is not None: + # Caller-provided size in BYTES (preferred for exact hardware OOB behavior). + if isinstance(num_records_bytes, int): + nbytes = int(num_records_bytes) + if nbytes <= 0: + nbytes = 0 + # Descriptor uses i32 bytes; clamp to the max representable. + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(nbytes) + else: + v = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(v.type, ir.IntegerType) or v.type.width != 64: + if isinstance(v.type, ir.IndexType): + op = std_arith.IndexCastOp(i64_type, v) + else: + op = std_arith.ExtSIOp(i64_type, v) + v = _unwrap_value(op.result) + num_records = v + elif max_size: + # Use max for flexibility (hardware will check actual bounds) + # Note: FlyDSL's rocdl.make.buffer.rsrc requires i32, not i64 + num_records = _create_i64_constant(0xFFFFFFFF) # FALLBACK_MAX_SIZE + else: + # Use the logical memref size (in bytes) for hardware OOB checking. + nbytes = _num_records_from_memref_type() + if nbytes is None: + # Fall back to max-size if we can't infer statically. + num_records = _create_i64_constant(0xFFFFFFFF) + else: + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(int(nbytes)) + + # Create resource descriptor (returns !llvm.ptr<8>) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + rsrc = rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride_val, num_records, flags).result + + return BufferResourceDescriptor(rsrc) + + +def create_buffer_resource_from_addr( + addr_i64: ir.Value, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from a raw i64 device address. + + Useful when working with runtime pointer arrays (e.g. IPC-mapped addresses + or device-side pointer tables) where no fly.memref is available. + The full address is encoded as the buffer base; callers should pass + byte offset 0 to buffer_load / buffer_store. + + Args: + addr_i64: Raw 64-bit device address (i64 MLIR value). + num_records_bytes: Optional buffer size in bytes for hardware OOB checking. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>). + + Example: + >>> rsrc = create_buffer_resource_from_addr(raw_addr_i64) + >>> data = buffer_load(rsrc, i32_zero, vec_width=4, dtype=T.i32) + """ + addr_i64 = _unwrap_value(addr_i64) + ptr_type = ir.Type.parse("!llvm.ptr") + base_ptr = llvm.IntToPtrOp(ptr_type, addr_i64).result + flags = _create_i32_constant(_get_buffer_flags()) + stride = _create_i16_constant(0) + if num_records_bytes is None: + num_records = _create_i64_constant(0xFFFFFFFF) + elif isinstance(num_records_bytes, int): + nbytes = max(0, min(int(num_records_bytes), 0xFFFFFFFF)) + num_records = _create_i64_constant(nbytes) + else: + num_records = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(num_records.type, ir.IntegerType) or num_records.type.width != 64: + if isinstance(num_records.type, ir.IndexType): + num_records = _unwrap_value(std_arith.IndexCastOp(i64_type, num_records).result) + else: + num_records = _unwrap_value(std_arith.ExtSIOp(i64_type, num_records).result) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + return rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride, num_records, flags).result + + +@traced_op +def create_buffer_resource( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from memref. + + This is a simplified wrapper around BufferResourceDescriptor.from_memref() + that returns the raw ROCDL resource value. + + Args: + memref_val: Memref value + stride: Buffer stride (0 for contiguous) + max_size: Use maximum buffer size + num_records_bytes: Override buffer size in bytes. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>) + + Example: + >>> rsrc = create_buffer_resource(A) + >>> data = buffer_load(rsrc, offset) + """ + desc = BufferResourceDescriptor.from_memref( + memref_val, + stride, + max_size, + num_records_bytes=num_records_bytes, + base_byte_offset=base_byte_offset, + ) + return desc.rsrc + + +@traced_op +def buffer_load( + rsrc: ir.Value, + offset: ir.Value, + vec_width: int = 4, + dtype=None, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + soffset_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """AMD buffer load operation. + + Load data from global memory using buffer descriptor and offset. + Uses hardware-level bounds checking and vectorization. + + Args: + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + vec_width: Vector width (1, 2, or 4) + dtype: Element data type (None for f32, or ir.F32Type, etc.) + mask: Optional mask for predicated load (i1 type) + cache_modifier: Cache control flags (0 for default) + soffset_bytes: Optional scalar offset (in BYTES) added by the buffer instruction (soffset). + Use this to fold small constant deltas into the instruction instead of emitting + extra VGPR address arithmetic. + + Returns: + Loaded data (scalar or vector depending on vec_width) + + Example: + >>> # Load 4xf32 + >>> data = buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Load with mask + >>> data = buffer_load(rsrc, offset, vec_width=4, mask=valid) + """ + # Default dtype to f32 + if dtype is None: + dtype = T.f32() + # Accept DSL Numeric class (e.g. fx.Int32) as dtype: unwrap to ir.Type + elif hasattr(dtype, "ir_type"): + dtype = dtype.ir_type + + # Unwrap offset first (accept Python ints and DSL Numeric values). + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: Buffer load offset is in BYTES, not elements! + # For vec4xf32, each element is 4 bytes, so multiply offset by 4 + element_bytes = dtype.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create vector type + if vec_width == 1: + result_type = dtype + else: + result_type = ir.VectorType.get([vec_width], dtype) + + # Create instruction offset and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(soffset_bytes) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer load + load_op = rocdl.RawPtrBufferLoadOp( + result_type, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) + + return load_op.result + + +@traced_op +def buffer_store( + data: ir.Value, + rsrc: ir.Value, + offset: ir.Value, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + *, + soffset_bytes: Optional[Union[int, ir.Value]] = None, + offset_is_bytes: bool = False, +): + """AMD buffer store operation. + + Store data to global memory using buffer descriptor and offset. + + Args: + data: Data to store (scalar or vector) + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + mask: Optional mask for predicated store (i1 type) + cache_modifier: Cache control flags (0 for default) + + Example: + >>> buffer_store(data, rsrc, offset) + >>> + >>> # Store with mask + >>> buffer_store(data, rsrc, offset, mask=valid) + """ + # Unwrap all inputs (accept DSL Numeric values via ir_value()) + if hasattr(data, "ir_value"): + data = data.ir_value() + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + data = _unwrap_value(data) + rsrc = _unwrap_value(rsrc) + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: RawPtrBufferStoreOp offset is in BYTES. + # For backward compat, `buffer_store()` accepts element offsets by default + # and scales them to bytes. Set `offset_is_bytes=True` to skip scaling. + if not offset_is_bytes: + # Get element size from data type + data_type = data.type + if hasattr(data_type, "element_type"): # Vector type + element_type = data_type.element_type + else: # Scalar type + element_type = data_type + element_bytes = element_type.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create instruction offset (soffset) and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(int(soffset_bytes)) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer store + rocdl.RawPtrBufferStoreOp( + data, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) diff --git a/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/meta.py b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/meta.py new file mode 100644 index 000000000..23e4c849a --- /dev/null +++ b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/meta.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +import inspect +from functools import wraps + +from flydsl._mlir import ir + + +def _to_raw_value(obj): + if isinstance(obj, ir.Value): + return obj + if isinstance(obj, type): + return obj + if hasattr(obj, "__extract_to_ir_values__"): + values = obj.__extract_to_ir_values__() + if len(values) != 1: + raise ValueError(f"Primitive function expects 1 value, got {len(values)}") + return values[0] + if isinstance(obj, tuple): + return tuple(_to_raw_value(e) for e in obj) + if isinstance(obj, list): + return [_to_raw_value(e) for e in obj] + return obj + + +def _flatten_args(args, kwargs): + new_args = tuple(_to_raw_value(a) for a in args) + new_kwargs = {k: _to_raw_value(v) if k not in ("loc", "ip") else v for k, v in kwargs.items()} + return new_args, new_kwargs + + +def _caller_location(depth=1): + """Build an MLIR Location from the Python call-site *depth* frames up.""" + frame = inspect.currentframe() + for _ in range(depth + 1): + if frame is not None: + frame = frame.f_back + if frame is None: + return ir.Location.unknown() + + info = inspect.getframeinfo(frame) + pos = getattr(info, "positions", None) + line = pos.lineno if pos is not None else info.lineno + col = (pos.col_offset or 0) if pos is not None else 0 + file_loc = ir.Location.file(info.filename, line, col) + + if info.code_context: + label = " ".join(ln.strip() for ln in info.code_context) + else: + label = info.function + return ir.Location.name(label, childLoc=file_loc) + + +def traced_op(op): + @wraps(op) + def wrapper(*args, **kwargs): + loc = kwargs.pop("loc", None) + if loc is None: + loc = _caller_location(depth=1) + args, kwargs = _flatten_args(args, kwargs) + with loc: + return op(*args, **kwargs) + + return wrapper diff --git a/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/vector.py b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/vector.py new file mode 100644 index 000000000..29f8289c2 --- /dev/null +++ b/tasks/torch2flydsl/moe_sorting_kernel/flydsl_compat/vector.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""Vector dialect helpers and re-exports. + +The ``Vector`` class itself lives in ``typing.py`` alongside other builtin +DSL types. This module re-exports it for convenience and provides thin +wrappers around upstream ``_mlir.dialects.vector`` ops. +""" + +from __future__ import annotations + +from flydsl._mlir import ir +from flydsl._mlir.dialects import vector as _vector + +# Re-export upstream dialect for ``from flydsl.expr import vector; vector.broadcast(...)`` +from flydsl._mlir.dialects.vector import * # noqa: F401,F403,E402 +from .meta import traced_op + +# Re-export Vector and friends so ``from flydsl.expr.vector import Vector`` works +from flydsl.expr.typing import ReductionOp, Vector, empty_like, full, full_like, ones_like, zeros_like # noqa: F401 + +# ═══════════════════════════════════════════════════════════════════════ +# Dialect helper wrappers (legacy, will be deprecated) +# Prefer using Vector methods or _mlir.dialects.vector directly. +# ═══════════════════════════════════════════════════════════════════════ + + +@traced_op +def from_elements(*args, loc=None, ip=None, **kwargs): + """Construct a vector from scalar elements, auto-unwrapping ArithValue wrappers.""" + from flydsl.expr import arith as _arith_ext + + if len(args) >= 2: + args = list(args) + elems = args[1] + if isinstance(elems, (list, tuple)): + args[1] = [_arith_ext.unwrap(v) for v in elems] + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + +@traced_op +def store(value, memref, indices, *, loc=None, ip=None, **kwargs): + """Vector store wrapper that accepts ArithValue/wrappers for value/indices.""" + from flydsl.expr import arith as _arith_ext + + return _vector.store( + _arith_ext.unwrap(value), + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + **kwargs, + ) + + +# ----------------------------------------------------------------------------- +# Thin wrappers for common op classes that otherwise require `.result` access. +# ----------------------------------------------------------------------------- + + +@traced_op +def extract(vector, static_position=None, dynamic_position=None, *, loc=None, ip=None): + """Wrapper around `vector.ExtractOp(...).result`. + + When only ``dynamic_position`` is supplied (without explicit + ``static_position``), each dynamic index needs a corresponding + ``kDynamic`` sentinel in the static attribute so the ODS builder + pairs them correctly. This wrapper fills in the sentinels + automatically. + """ + from flydsl.expr import arith as _arith_ext + + if static_position is None: + static_position = [] + if dynamic_position is None: + dynamic_position = [] + dynamic_position = [_arith_ext.unwrap(i, index=True) for i in dynamic_position] + + n_static = len(static_position) + n_dynamic = len(dynamic_position) + if n_dynamic > 0 and n_static < n_dynamic: + kDynamic = ir.ShapedType.get_dynamic_size() + static_position = list(static_position) + [kDynamic] * (n_dynamic - n_static) + + return _vector.ExtractOp( + _arith_ext.unwrap(vector), + static_position=static_position, + dynamic_position=dynamic_position, + loc=loc, + ip=ip, + ).result + + +@traced_op +def load_op(result_type, memref, indices, *, loc=None, ip=None): + """Wrapper around `vector.LoadOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.LoadOp( + result_type, + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + ).result + + +@traced_op +def bitcast(result_type, source, *, loc=None, ip=None): + """Wrapper around `vector.BitCastOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.BitCastOp( + result_type, + _arith_ext.unwrap(source), + loc=loc, + ip=ip, + ).result diff --git a/tasks/torch2flydsl/moe_sorting_kernel/kernel.py b/tasks/torch2flydsl/moe_sorting_kernel/kernel.py index 893c154f6..46c6fe1c2 100644 --- a/tasks/torch2flydsl/moe_sorting_kernel/kernel.py +++ b/tasks/torch2flydsl/moe_sorting_kernel/kernel.py @@ -27,7 +27,8 @@ from flydsl._mlir import ir from flydsl._mlir.dialects import memref as memref_ops from flydsl.compiler.kernel_function import CompilationContext -from flydsl.expr import buffer_ops, gpu, range_constexpr +from flydsl.expr import gpu, range_constexpr +from flydsl_compat import buffer_ops from flydsl.expr import rocdl as fly_rocdl from flydsl.expr.arith import ArithValue from flydsl.expr.typing import T diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/README.md b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/README.md index eb041a440..1dd2e2e5f 100644 --- a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/README.md +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/README.md @@ -1,6 +1,10 @@ # qk_norm_rope_quant_kernel: task-owned v2 contract -Implement or optimize qk_norm_rope_quant in FlyDSL, preserving all task inputs, outputs and numerical gates. +Implement or optimize the BF16 RMSNorm + RoPE (`quant=False`) path of +`flydsl_qk_norm_rope_quant` in FlyDSL, preserving all six declared workloads, +outputs and numerical gates. Optional FP8 quantization (`quant=True`) and its +packed codes/scales are outside this task's evaluated contract; these cases +do not qualify that separate capability. The candidate starts **implemented**. Arena freezes the implemented source in a separate baseline workspace. The original harness's primary implementation timing is retained; additional @@ -17,9 +21,8 @@ reference against the original independent AITER comparison wherever that check was present. Small independent known-answer and negative-output controls in `scripts/reference_controls.py` supplement the full GPU checks. -Edit only `candidate.editable` paths from `config.yaml`. Preserve each declared -public operator/builder interface and all outputs (including residuals, packed -quantization codes/scales, routing indices or state when applicable). Inspect the +Edit only `candidate.editable` paths from `config.yaml`. Preserve the declared entrypoint and the evaluated four-element return tuple: +BF16 Q, BF16 KV, None and None. Inspect the protected harness calls and `model.py` to understand shapes, strides and layout. The final operator computation must run FlyDSL GPU kernels. PyTorch is allowed for allocation, views and launch preparation, not replacement operator compute. @@ -71,11 +74,12 @@ outside timing, poison both output tensors and replay the same measured invocation against the unchanged model. Recheck the two None scale slots on replay too, and restore the inputs. Baseline/candidate retain the original 10external warmups,100samples and graph timing; diagnostic model timing retains -its original10warmups. Source kernel, model, cases, group-size variants and seed +its original10warmups. Kernel computation, model, cases, group-size variants and seed are unchanged. Final candidate arithmetic must run FlyDSL; the candidate-only auditor permits host preparation and checks operator calls outside timing. -This original source uses the older FlyDSL buffer_ops API, so its full GPU -qualification requires the pinned compatible image recorded with the report. +The legacy buffer/vector calls use the bundled compatibility adapters when +the installed runtime removes those helpers. Full GPU qualification must bind +the task source and the selected image. The allowed preparation dependencies include exactly `from aiter.utility import dtypes` (with an optional alias), used by the original @@ -98,3 +102,8 @@ passing the module to another function, dynamic attribute access, other members (including any library imported by the dtype module), and relative/package import alternatives are rejected. Reading a hardware dtype constant does not expose an AITER operator dependency to the candidate. + +The task-local `flydsl_compat` helpers preserve the legacy buffer/vector API +when the installed FlyDSL no longer supplies it. See `flydsl_compat/SOURCE.md` +for the pinned upstream source and retained license. Workloads, numerical gates, +and timing parameters are unchanged. diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/config.yaml b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/config.yaml index 65af2f176..2f7e64f5d 100644 --- a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/config.yaml +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/config.yaml @@ -1,6 +1,6 @@ schema_version: 2 -description: Implement or optimize qk_norm_rope_quant in FlyDSL, preserving all task inputs, outputs and - numerical gates. +description: Implement or optimize the six BF16 quant=False RMSNorm + RoPE workloads in FlyDSL, + preserving Q/KV outputs, None scale slots and numerical gates. Optional quant=True FP8 output is not evaluated. instructions: - README.md candidate: diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/LICENSE b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/LICENSE new file mode 100644 index 000000000..c73e2f2c4 --- /dev/null +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/LICENSE @@ -0,0 +1,17 @@ + Apache License + Version 2.0, January 2004 + http://www.apache.org/licenses/ + + Copyright 2025 FlyDSL Project Contributors + + Licensed under the Apache License, Version 2.0 (the "License"); + you may not use this file except in compliance with the License. + You may obtain a copy of the License at + + http://www.apache.org/licenses/LICENSE-2.0 + + Unless required by applicable law or agreed to in writing, software + distributed under the License is distributed on an "AS IS" BASIS, + WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. + See the License for the specific language governing permissions and + limitations under the License. diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/SOURCE.md b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/SOURCE.md new file mode 100644 index 000000000..83ea5ad77 --- /dev/null +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/SOURCE.md @@ -0,0 +1,13 @@ +# Pinned FlyDSL compatibility helpers + +Derived from ROCm/FlyDSL commit `28a18d328b4882c999864b2df2f8f9fe3fcc8b47`, +`python/flydsl/expr/{buffer_ops,vector,meta}.py`, under Apache-2.0 (see LICENSE). +The original buffer descriptor, byte offset, masking and cache policy is retained. +Relative dependency imports now target the installed package. Removed memref +pointer extraction uses the current typed iterator and LLVM pointer conversion. +Legacy runtimes continue using their installed original helpers. No runtime +monkeypatching, external repository imports or downloads are performed. + +Current ROCDL load/store operations receive the same cache-policy bits through +their `aux` attribute instead of a removed SSA operand. Vector unwrapping uses +the current signature while retaining the surrounding MLIR location context. diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/__init__.py b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/__init__.py new file mode 100644 index 000000000..1bbff9ffa --- /dev/null +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/__init__.py @@ -0,0 +1,13 @@ +"""Task-local legacy API fallback; never modifies the installed FlyDSL package.""" +try: + import flydsl.expr.buffer_ops as buffer_ops +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.buffer_ops": + raise + from . import buffer_ops +try: + import flydsl.expr.vector as vector +except ModuleNotFoundError as exc: + if exc.name != "flydsl.expr.vector": + raise + from . import vector diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/buffer_ops.py b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/buffer_ops.py new file mode 100644 index 000000000..ab27dd821 --- /dev/null +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/buffer_ops.py @@ -0,0 +1,603 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""AMD Buffer Load/Store Operations - High-level Python API + +This module provides high-level Python wrappers for AMD CDNA3/CDNA4 buffer operations. +Buffer operations use a scalar base pointer and per-thread offsets for efficient memory access. + +Example: + >>> from flydsl._mlir_helpers import buffer_ops + >>> from flydsl._mlir_helpers import arith + >>> import _mlir.extras.types as T + >>> + >>> # Create buffer resource from memref + >>> rsrc = buffer_ops.create_buffer_resource(A) + >>> + >>> # Compute offset + >>> offset = row * arith.index(4096) + col + >>> + >>> # Buffer load (4xf32) + >>> data = buffer_ops.buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Buffer store + >>> buffer_ops.buffer_store(data, rsrc, offset) +""" + +from typing import Optional, Union + +from flydsl._mlir import ir +from flydsl._mlir.dialects import arith as std_arith +from flydsl._mlir.dialects import llvm, rocdl +from flydsl._mlir.extras import types as T +from flydsl.runtime.device import is_rdna_arch +from .meta import traced_op + + +def _get_buffer_flags(arch=None): + """Get AMD buffer resource descriptor (V#) flags word (bits 127:96). + + Constructs the 32-bit flags field for rocdl.make.buffer.rsrc, following the + same logic as LLVM's AMDGPUToROCDL makeBufferRsrc(): + https://github.com/llvm/llvm-project/blob/main/mlir/lib/Conversion/AMDGPUToROCDL/AMDGPUToROCDL.cpp + + Bit layout (common to all architectures): + bits [11:0] - DST_SEL: ignored by raw buffer intrinsics + bits [14:12] - DATA_FORMAT: must be nonzero, 7 = float + bits [18:15] - NUM_FORMAT: must be nonzero, 4 = 32-bit + bit [19] - In nested heap (0) + bit [20] - Behavior on unmap (0 = return 0 / ignore) + bits [22:21] - Index stride for swizzles (0) + bit [23] - Add thread ID (0) + bit [24] - Reserved: must be 1 on RDNA, 0 on CDNA + bits [26:25] - Reserved (0) + bit [27] - Non-volatile (CDNA only, 0) + bits [29:28] - OOB_SELECT (RDNA only): 0=structured, 2=none, 3=check offset + bits [31:30] - Type (must be 0) + + CDNA (gfx9xx): (7 << 12) | (4 << 15) = 0x20070 + RDNA (gfx10+): (7 << 12) | (4 << 15) | (1 << 24) | (2 << 28) = 0x21020070 + - bit 24 set to 1 (required on RDNA) + - OOB_SELECT=2 (no bounds checking, matching LLVM boundsCheck=false) + """ + import os + + if arch is None: + arch = os.environ.get("FLYDSL_GPU_ARCH") + flags = (7 << 12) | (4 << 15) + if is_rdna_arch(arch): + flags |= 1 << 24 # reserved bit, must be 1 on RDNA + flags |= 2 << 28 # OOB_SELECT = 2 (no bounds checking) + return flags + + +__all__ = [ + "create_llvm_ptr", + "get_element_ptr", + "create_buffer_resource", + "create_buffer_resource_from_addr", + "buffer_load", + "buffer_store", + "BufferResourceDescriptor", + "extract_base_index", +] + + +def _unwrap_value(value): + """Recursively unwrap ArithValue or similar wrappers to get the actual MLIR value. + + Handles: + - FlyDSL ArithValue (has ._value) + - flyc DSL Numeric like fx.Int32 (has .ir_value() method) + - flyc ArithValue (is already ir.Value subclass) + """ + # DSL Numeric (Int32, Float32, etc.) — use ir_value() to materialize + if hasattr(value, "ir_value") and not isinstance(value, ir.Value): + return value.ir_value() + max_depth = 10 # Safety limit + depth = 0 + while depth < max_depth and not isinstance(value, ir.Value): + if hasattr(value, "_value"): + value = value._value + elif hasattr(value, "value"): + value = value.value + else: + break + depth += 1 + return value + + +def _create_i32_constant(value: int) -> ir.Value: + """Create i32 constant using standard MLIR arith dialect.""" + i32_type = T.i32() + if value > 0x7FFFFFFF: + value = int(value - 2**32) + attr = ir.IntegerAttr.get(i32_type, value) + op = std_arith.ConstantOp(i32_type, attr) + return _unwrap_value(op.result) + + +def _create_i16_constant(value: int) -> ir.Value: + """Create i16 constant using standard MLIR arith dialect.""" + i16_type = T.i16() + attr = ir.IntegerAttr.get(i16_type, value) + op = std_arith.ConstantOp(i16_type, attr) + return _unwrap_value(op.result) + + +def _create_i64_constant(value: int) -> ir.Value: + """Create i64 constant using standard MLIR arith dialect.""" + i64_type = T.i64() + attr = ir.IntegerAttr.get(i64_type, value) + op = std_arith.ConstantOp(i64_type, attr) + return _unwrap_value(op.result) + + +def create_llvm_ptr(value, address_space: int = 0) -> ir.Value: + """Create an LLVM pointer from an integer or index value.""" + value = _unwrap_value(value) + if isinstance(value.type, ir.IndexType): + i64_type = T.i64() + value = _unwrap_value(std_arith.IndexCastOp(i64_type, value).result) + ptr_type = ir.Type.parse(f"!llvm.ptr<{address_space}>") + return llvm.IntToPtrOp(ptr_type, value).result + + +def extract_base_index(tensor, address_space: int = 1) -> ir.Value: + """Extract the base address of a fly.memref as an index value. + + Inverse of :func:`create_llvm_ptr` (index -> ptr). Useful when ISA + requires a raw pointer instead of a buffer resource descriptor + (e.g. global_atomic_pk_add_bf16 on gfx942). + """ + from flydsl._mlir.dialects import fly as _fly + from flydsl._mlir.dialects import memref as _memref + + raw = _unwrap_value(tensor) + try: + ir.MemRefType(raw.type) + return _memref.extract_aligned_pointer_as_index(raw) + except ValueError: + pass + + # FlyDSL 0.3 uses a typed iterator instead of the removed extract op. + from flydsl.expr import get_iter, to_llvm_ptr + ptr = to_llvm_ptr(get_iter(raw)) + i64_val = llvm.PtrToIntOp(ir.IntegerType.get_signless(64), ptr).result + return _unwrap_value(std_arith.IndexCastOp(ir.IndexType.get(), i64_val).result) + + +def get_element_ptr( + base_ptr, + byte_offset: Union[int, ir.Value, None] = None, + static_byte_offset: int = 0, + elem_type: Optional[ir.Type] = None, + no_wrap_flags=None, +) -> ir.Value: + """Build an LLVM GEP from a base pointer plus byte offsets.""" + _gep_dynamic_index_sentinel = -(2**31) + + base_ptr = _unwrap_value(base_ptr) + if not isinstance(static_byte_offset, int): + raise TypeError(f"static_byte_offset must be int, got {type(static_byte_offset).__name__}") + if elem_type is None: + elem_type = T.i8() + elif callable(elem_type): + elem_type = elem_type() + + if byte_offset is None: + dynamic_indices = [] + raw_constant_indices = [int(static_byte_offset)] + elif isinstance(byte_offset, int): + dynamic_indices = [] + raw_constant_indices = [int(byte_offset) + int(static_byte_offset)] + else: + offset_val = _unwrap_value(byte_offset) + if isinstance(offset_val.type, ir.IndexType): + i64_type = T.i64() + offset_val = _unwrap_value(std_arith.IndexCastOp(i64_type, offset_val).result) + elif not isinstance(offset_val.type, ir.IntegerType): + raise TypeError("byte_offset must be int, index, or integer-typed MLIR value; " f"got {offset_val.type}") + + if static_byte_offset != 0: + static_type = offset_val.type + static_attr = ir.IntegerAttr.get(static_type, int(static_byte_offset)) + static_const = _unwrap_value(std_arith.ConstantOp(static_type, static_attr).result) + offset_val = _unwrap_value(std_arith.AddIOp(offset_val, static_const).result) + + dynamic_indices = [offset_val] + raw_constant_indices = [_gep_dynamic_index_sentinel] + + return llvm.GEPOp( + base_ptr.type, + base_ptr, + dynamic_indices, + raw_constant_indices, + elem_type, + no_wrap_flags, + ).result + + +class BufferResourceDescriptor: + """AMD Buffer Resource Descriptor + + A buffer resource descriptor contains: + - base_pointer: Scalar base pointer (wave-uniform, stored in SGPRs) + - stride: Stride for structured buffers (typically 0 for contiguous) + - num_records: Buffer size in bytes + - flags: Data format and access flags + + The descriptor is stored in a special LLVM pointer type (!llvm.ptr<8>) + """ + + def __init__(self, rsrc: ir.Value): + """Initialize with ROCDL resource descriptor value.""" + self.rsrc = rsrc + + @staticmethod + def from_memref( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + data_format: str = "f32", + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, + ) -> "BufferResourceDescriptor": + """Create buffer resource descriptor from memref. + + Args: + memref_val: Memref value to create descriptor for + stride: Stride in elements (0 for contiguous) + max_size: If True, use max buffer size for flexibility + num_records_bytes: Override buffer size (in BYTES) used by hardware OOB checking. + If provided, this takes precedence over `max_size`. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + data_format: Data format ('f32', 'f16', 'i32', etc.) + + Returns: + BufferResourceDescriptor instance + + Example: + >>> rsrc = BufferResourceDescriptor.from_memref(A) + """ + # Extract raw pointer from fly.memref. + raw_val = _unwrap_value(memref_val) + from flydsl._mlir.dialects import fly as _fly + + # Preserve byte offsets and descriptor policy with the current pointer API. + from flydsl.expr import get_iter, to_llvm_ptr + base_ptr = to_llvm_ptr(get_iter(raw_val)) + if base_byte_offset is not None: + base_ptr = get_element_ptr(base_ptr, byte_offset=base_byte_offset) + + # Create buffer resource descriptor + flags_val = _get_buffer_flags() + flags = _create_i32_constant(flags_val) + stride_val = _create_i16_constant(stride) + + def _num_records_from_memref_type() -> Optional[int]: + """Best-effort: derive logical buffer size (in bytes) from static memref type.""" + try: + mt = ir.MemRefType(_unwrap_value(memref_val).type) + shape = list(mt.shape) + if any(int(d) < 0 for d in shape): + return None + # Compute element size in bytes (scalar element type). + elem_t = mt.element_type + elem_bits = getattr(elem_t, "width", None) + if elem_bits is None: + return None + elem_bytes = int(elem_bits) // 8 + if elem_bytes <= 0: + return None + num_elems = 1 + for d in shape: + num_elems *= int(d) + return int(num_elems) * int(elem_bytes) + except Exception: + return None + + if num_records_bytes is not None: + # Caller-provided size in BYTES (preferred for exact hardware OOB behavior). + if isinstance(num_records_bytes, int): + nbytes = int(num_records_bytes) + if nbytes <= 0: + nbytes = 0 + # Descriptor uses i32 bytes; clamp to the max representable. + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(nbytes) + else: + v = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(v.type, ir.IntegerType) or v.type.width != 64: + if isinstance(v.type, ir.IndexType): + op = std_arith.IndexCastOp(i64_type, v) + else: + op = std_arith.ExtSIOp(i64_type, v) + v = _unwrap_value(op.result) + num_records = v + elif max_size: + # Use max for flexibility (hardware will check actual bounds) + # Note: FlyDSL's rocdl.make.buffer.rsrc requires i32, not i64 + num_records = _create_i64_constant(0xFFFFFFFF) # FALLBACK_MAX_SIZE + else: + # Use the logical memref size (in bytes) for hardware OOB checking. + nbytes = _num_records_from_memref_type() + if nbytes is None: + # Fall back to max-size if we can't infer statically. + num_records = _create_i64_constant(0xFFFFFFFF) + else: + if nbytes > 0xFFFFFFFF: + nbytes = 0xFFFFFFFF + num_records = _create_i64_constant(int(nbytes)) + + # Create resource descriptor (returns !llvm.ptr<8>) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + rsrc = rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride_val, num_records, flags).result + + return BufferResourceDescriptor(rsrc) + + +def create_buffer_resource_from_addr( + addr_i64: ir.Value, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from a raw i64 device address. + + Useful when working with runtime pointer arrays (e.g. IPC-mapped addresses + or device-side pointer tables) where no fly.memref is available. + The full address is encoded as the buffer base; callers should pass + byte offset 0 to buffer_load / buffer_store. + + Args: + addr_i64: Raw 64-bit device address (i64 MLIR value). + num_records_bytes: Optional buffer size in bytes for hardware OOB checking. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>). + + Example: + >>> rsrc = create_buffer_resource_from_addr(raw_addr_i64) + >>> data = buffer_load(rsrc, i32_zero, vec_width=4, dtype=T.i32) + """ + addr_i64 = _unwrap_value(addr_i64) + ptr_type = ir.Type.parse("!llvm.ptr") + base_ptr = llvm.IntToPtrOp(ptr_type, addr_i64).result + flags = _create_i32_constant(_get_buffer_flags()) + stride = _create_i16_constant(0) + if num_records_bytes is None: + num_records = _create_i64_constant(0xFFFFFFFF) + elif isinstance(num_records_bytes, int): + nbytes = max(0, min(int(num_records_bytes), 0xFFFFFFFF)) + num_records = _create_i64_constant(nbytes) + else: + num_records = _unwrap_value(num_records_bytes) + i64_type = T.i64() + if not isinstance(num_records.type, ir.IntegerType) or num_records.type.width != 64: + if isinstance(num_records.type, ir.IndexType): + num_records = _unwrap_value(std_arith.IndexCastOp(i64_type, num_records).result) + else: + num_records = _unwrap_value(std_arith.ExtSIOp(i64_type, num_records).result) + rsrc_type = ir.Type.parse("!llvm.ptr<8>") + return rocdl.MakeBufferRsrcOp(rsrc_type, base_ptr, stride, num_records, flags).result + + +@traced_op +def create_buffer_resource( + memref_val: ir.Value, + stride: int = 0, + max_size: bool = True, + *, + num_records_bytes: Optional[Union[int, ir.Value]] = None, + base_byte_offset: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """Create AMD buffer resource descriptor from memref. + + This is a simplified wrapper around BufferResourceDescriptor.from_memref() + that returns the raw ROCDL resource value. + + Args: + memref_val: Memref value + stride: Buffer stride (0 for contiguous) + max_size: Use maximum buffer size + num_records_bytes: Override buffer size in bytes. + base_byte_offset: Optional byte offset added to the descriptor base pointer. + + Returns: + ROCDL buffer resource descriptor (!llvm.ptr<8>) + + Example: + >>> rsrc = create_buffer_resource(A) + >>> data = buffer_load(rsrc, offset) + """ + desc = BufferResourceDescriptor.from_memref( + memref_val, + stride, + max_size, + num_records_bytes=num_records_bytes, + base_byte_offset=base_byte_offset, + ) + return desc.rsrc + + +@traced_op +def buffer_load( + rsrc: ir.Value, + offset: ir.Value, + vec_width: int = 4, + dtype=None, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + soffset_bytes: Optional[Union[int, ir.Value]] = None, +) -> ir.Value: + """AMD buffer load operation. + + Load data from global memory using buffer descriptor and offset. + Uses hardware-level bounds checking and vectorization. + + Args: + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + vec_width: Vector width (1, 2, or 4) + dtype: Element data type (None for f32, or ir.F32Type, etc.) + mask: Optional mask for predicated load (i1 type) + cache_modifier: Cache control flags (0 for default) + soffset_bytes: Optional scalar offset (in BYTES) added by the buffer instruction (soffset). + Use this to fold small constant deltas into the instruction instead of emitting + extra VGPR address arithmetic. + + Returns: + Loaded data (scalar or vector depending on vec_width) + + Example: + >>> # Load 4xf32 + >>> data = buffer_load(rsrc, offset, vec_width=4) + >>> + >>> # Load with mask + >>> data = buffer_load(rsrc, offset, vec_width=4, mask=valid) + """ + # Default dtype to f32 + if dtype is None: + dtype = T.f32() + # Accept DSL Numeric class (e.g. fx.Int32) as dtype: unwrap to ir.Type + elif hasattr(dtype, "ir_type"): + dtype = dtype.ir_type + + # Unwrap offset first (accept Python ints and DSL Numeric values). + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: Buffer load offset is in BYTES, not elements! + # For vec4xf32, each element is 4 bytes, so multiply offset by 4 + element_bytes = dtype.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create vector type + if vec_width == 1: + result_type = dtype + else: + result_type = ir.VectorType.get([vec_width], dtype) + + # Create instruction offset and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(soffset_bytes) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer load + load_op = rocdl.RawPtrBufferLoadOp( + result_type, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) + + return load_op.result + + +@traced_op +def buffer_store( + data: ir.Value, + rsrc: ir.Value, + offset: ir.Value, + mask: Optional[ir.Value] = None, + cache_modifier: int = 0, + *, + soffset_bytes: Optional[Union[int, ir.Value]] = None, + offset_is_bytes: bool = False, +): + """AMD buffer store operation. + + Store data to global memory using buffer descriptor and offset. + + Args: + data: Data to store (scalar or vector) + rsrc: Buffer resource descriptor (!llvm.ptr<8>) + offset: Offset in elements (i32 type) + mask: Optional mask for predicated store (i1 type) + cache_modifier: Cache control flags (0 for default) + + Example: + >>> buffer_store(data, rsrc, offset) + >>> + >>> # Store with mask + >>> buffer_store(data, rsrc, offset, mask=valid) + """ + # Unwrap all inputs (accept DSL Numeric values via ir_value()) + if hasattr(data, "ir_value"): + data = data.ir_value() + if isinstance(offset, int): + offset = _create_i32_constant(offset) + elif hasattr(offset, "ir_value"): + offset = offset.ir_value() + data = _unwrap_value(data) + rsrc = _unwrap_value(rsrc) + offset = _unwrap_value(offset) + + # Convert offset to i32 if needed + if not isinstance(offset.type, ir.IntegerType) or offset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), offset) + offset = _unwrap_value(op.result) + + # IMPORTANT: RawPtrBufferStoreOp offset is in BYTES. + # For backward compat, `buffer_store()` accepts element offsets by default + # and scales them to bytes. Set `offset_is_bytes=True` to skip scaling. + if not offset_is_bytes: + # Get element size from data type + data_type = data.type + if hasattr(data_type, "element_type"): # Vector type + element_type = data_type.element_type + else: # Scalar type + element_type = data_type + element_bytes = element_type.width // 8 + bytes_const = _create_i32_constant(element_bytes) + op = std_arith.MulIOp(offset, bytes_const) + offset = _unwrap_value(op.result) + + # Apply mask by setting invalid offsets to max + if mask is not None: + mask = _unwrap_value(mask) + max_offset = _create_i32_constant(0x7FFFFFFF) + op = std_arith.SelectOp(mask, offset, max_offset) + offset = _unwrap_value(op.result) + + # Create instruction offset (soffset) and aux flags + if soffset_bytes is None: + soffset = _create_i32_constant(0) + else: + if isinstance(soffset_bytes, int): + soffset = _create_i32_constant(int(soffset_bytes)) + else: + soffset = _unwrap_value(soffset_bytes) + if not isinstance(soffset.type, ir.IntegerType) or soffset.type.width != 32: + op = std_arith.IndexCastOp(T.i32(), soffset) + soffset = _unwrap_value(op.result) + aux_flags = ir.IntegerAttr.get(T.i32(), cache_modifier) + + # Emit buffer store + rocdl.RawPtrBufferStoreOp( + data, rsrc, offset, soffset, aux=aux_flags # soffset (scalar byte offset) # aux (cache modifiers) + ) diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/meta.py b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/meta.py new file mode 100644 index 000000000..23e4c849a --- /dev/null +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/meta.py @@ -0,0 +1,65 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +import inspect +from functools import wraps + +from flydsl._mlir import ir + + +def _to_raw_value(obj): + if isinstance(obj, ir.Value): + return obj + if isinstance(obj, type): + return obj + if hasattr(obj, "__extract_to_ir_values__"): + values = obj.__extract_to_ir_values__() + if len(values) != 1: + raise ValueError(f"Primitive function expects 1 value, got {len(values)}") + return values[0] + if isinstance(obj, tuple): + return tuple(_to_raw_value(e) for e in obj) + if isinstance(obj, list): + return [_to_raw_value(e) for e in obj] + return obj + + +def _flatten_args(args, kwargs): + new_args = tuple(_to_raw_value(a) for a in args) + new_kwargs = {k: _to_raw_value(v) if k not in ("loc", "ip") else v for k, v in kwargs.items()} + return new_args, new_kwargs + + +def _caller_location(depth=1): + """Build an MLIR Location from the Python call-site *depth* frames up.""" + frame = inspect.currentframe() + for _ in range(depth + 1): + if frame is not None: + frame = frame.f_back + if frame is None: + return ir.Location.unknown() + + info = inspect.getframeinfo(frame) + pos = getattr(info, "positions", None) + line = pos.lineno if pos is not None else info.lineno + col = (pos.col_offset or 0) if pos is not None else 0 + file_loc = ir.Location.file(info.filename, line, col) + + if info.code_context: + label = " ".join(ln.strip() for ln in info.code_context) + else: + label = info.function + return ir.Location.name(label, childLoc=file_loc) + + +def traced_op(op): + @wraps(op) + def wrapper(*args, **kwargs): + loc = kwargs.pop("loc", None) + if loc is None: + loc = _caller_location(depth=1) + args, kwargs = _flatten_args(args, kwargs) + with loc: + return op(*args, **kwargs) + + return wrapper diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/vector.py b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/vector.py new file mode 100644 index 000000000..29f8289c2 --- /dev/null +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/flydsl_compat/vector.py @@ -0,0 +1,121 @@ +# SPDX-License-Identifier: Apache-2.0 +# Copyright (c) 2025 FlyDSL Project Contributors + +"""Vector dialect helpers and re-exports. + +The ``Vector`` class itself lives in ``typing.py`` alongside other builtin +DSL types. This module re-exports it for convenience and provides thin +wrappers around upstream ``_mlir.dialects.vector`` ops. +""" + +from __future__ import annotations + +from flydsl._mlir import ir +from flydsl._mlir.dialects import vector as _vector + +# Re-export upstream dialect for ``from flydsl.expr import vector; vector.broadcast(...)`` +from flydsl._mlir.dialects.vector import * # noqa: F401,F403,E402 +from .meta import traced_op + +# Re-export Vector and friends so ``from flydsl.expr.vector import Vector`` works +from flydsl.expr.typing import ReductionOp, Vector, empty_like, full, full_like, ones_like, zeros_like # noqa: F401 + +# ═══════════════════════════════════════════════════════════════════════ +# Dialect helper wrappers (legacy, will be deprecated) +# Prefer using Vector methods or _mlir.dialects.vector directly. +# ═══════════════════════════════════════════════════════════════════════ + + +@traced_op +def from_elements(*args, loc=None, ip=None, **kwargs): + """Construct a vector from scalar elements, auto-unwrapping ArithValue wrappers.""" + from flydsl.expr import arith as _arith_ext + + if len(args) >= 2: + args = list(args) + elems = args[1] + if isinstance(elems, (list, tuple)): + args[1] = [_arith_ext.unwrap(v) for v in elems] + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + return _vector.from_elements(*args, loc=loc, ip=ip, **kwargs) + + +@traced_op +def store(value, memref, indices, *, loc=None, ip=None, **kwargs): + """Vector store wrapper that accepts ArithValue/wrappers for value/indices.""" + from flydsl.expr import arith as _arith_ext + + return _vector.store( + _arith_ext.unwrap(value), + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + **kwargs, + ) + + +# ----------------------------------------------------------------------------- +# Thin wrappers for common op classes that otherwise require `.result` access. +# ----------------------------------------------------------------------------- + + +@traced_op +def extract(vector, static_position=None, dynamic_position=None, *, loc=None, ip=None): + """Wrapper around `vector.ExtractOp(...).result`. + + When only ``dynamic_position`` is supplied (without explicit + ``static_position``), each dynamic index needs a corresponding + ``kDynamic`` sentinel in the static attribute so the ODS builder + pairs them correctly. This wrapper fills in the sentinels + automatically. + """ + from flydsl.expr import arith as _arith_ext + + if static_position is None: + static_position = [] + if dynamic_position is None: + dynamic_position = [] + dynamic_position = [_arith_ext.unwrap(i, index=True) for i in dynamic_position] + + n_static = len(static_position) + n_dynamic = len(dynamic_position) + if n_dynamic > 0 and n_static < n_dynamic: + kDynamic = ir.ShapedType.get_dynamic_size() + static_position = list(static_position) + [kDynamic] * (n_dynamic - n_static) + + return _vector.ExtractOp( + _arith_ext.unwrap(vector), + static_position=static_position, + dynamic_position=dynamic_position, + loc=loc, + ip=ip, + ).result + + +@traced_op +def load_op(result_type, memref, indices, *, loc=None, ip=None): + """Wrapper around `vector.LoadOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.LoadOp( + result_type, + _arith_ext.unwrap(memref), + [_arith_ext.unwrap(i, index=True) for i in indices], + loc=loc, + ip=ip, + ).result + + +@traced_op +def bitcast(result_type, source, *, loc=None, ip=None): + """Wrapper around `vector.BitCastOp(...).result`.""" + from flydsl.expr import arith as _arith_ext + + return _vector.BitCastOp( + result_type, + _arith_ext.unwrap(source), + loc=loc, + ip=ip, + ).result diff --git a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/kernel.py b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/kernel.py index ac3a0d5d5..5e8e72769 100644 --- a/tasks/torch2flydsl/qk_norm_rope_quant_kernel/kernel.py +++ b/tasks/torch2flydsl/qk_norm_rope_quant_kernel/kernel.py @@ -35,11 +35,12 @@ import flydsl.compiler as flyc import flydsl.expr as fx -from flydsl.expr import arith, buffer_ops, const_expr, ptrtoint, range_constexpr, vector +from flydsl.expr import arith, const_expr, ptrtoint, range_constexpr +from flydsl_compat import buffer_ops, vector from flydsl.expr import math as fmath from flydsl.expr.arith import ArithValue, CmpFPredicate from flydsl.expr.typing import Int32, Stream, T -from flydsl.expr.vector import ReductionOp +ReductionOp = vector.ReductionOp from flydsl._mlir import ir from flydsl._mlir.dialects import fly, llvm, rocdl from flydsl.compiler.protocol import extract_to_ir_values diff --git a/tasks/triton2flydsl/aiter/mla/README.md b/tasks/triton2flydsl/aiter/mla/README.md index 1763c2d1a..efe58f1c9 100644 --- a/tasks/triton2flydsl/aiter/mla/README.md +++ b/tasks/triton2flydsl/aiter/mla/README.md @@ -67,9 +67,9 @@ and the attention output negate. Page addresses and lengths stay unchanged. Output poisoning, independent reference work and restoration are outside timing. All original seeds42+i, ten external warmups and100graph samples remain. -Runtime qualification must bind the exact image. The unchanged source's -multi-stage pipeline has shown non-finite output for the128-head/lora512case in -a newer Triton runtime; the original pinned runtime passed all six diagnostic -cases. Diagnostics alone do not qualify the task. Use a full validator report -for the selected runtime; do not waive that case, alter its tolerance, or select -a different baseline implementation after a session has frozen it. +Runtime qualification must bind the exact image. The gfx950 decode kernel uses +one pipeline stage: the two-stage schedule produced incorrect and non-finite +output on the ROCm 10 Triton runtime. The operator, all six workloads and their +numerical gates remain unchanged. Repeated GPU checks include additional seeds +for the 128-head/lora512 case. A full validator report is still required for the +selected runtime; do not change the baseline after a session has frozen it. diff --git a/tasks/triton2flydsl/aiter/mla/mla.py b/tasks/triton2flydsl/aiter/mla/mla.py index 970807319..c72f6ec52 100644 --- a/tasks/triton2flydsl/aiter/mla/mla.py +++ b/tasks/triton2flydsl/aiter/mla/mla.py @@ -924,7 +924,9 @@ def select_3d_config( "NUM_SEGMENTS_PER_SEQ": num_segments, "num_warps": attn_num_warps, "waves_per_eu": attn_waves_per_eu, - "num_stages": 2 if DEVICE_ARCH in ("gfx1250", "gfx950") else 1, + # The two-stage gfx950 pipeline produces incorrect values with Triton 3.8. + # Keep the single-stage schedule; workload and numerical gates are unchanged. + "num_stages": 2 if DEVICE_ARCH == "gfx1250" else 1, } reduce_config = { "TILE_SIZE": TILE_SIZE, diff --git a/tests/eval_tools/test_manager.py b/tests/eval_tools/test_manager.py index e2783ad52..508052cea 100644 --- a/tests/eval_tools/test_manager.py +++ b/tests/eval_tools/test_manager.py @@ -3,6 +3,8 @@ from dataclasses import dataclass, replace from pathlib import Path +import pytest + from src.eval_tools.config import EvalToolsConfig from src.eval_tools.contracts import ( CapabilityCheck, @@ -169,9 +171,14 @@ def test_repeated_evaluation_cannot_reuse_stale_tool_artifacts(tmp_path): assert not (second_dir / "build_attestation.json").exists() -def test_runtime_evidence_cannot_be_shadowed_by_config_options(tmp_path): +@pytest.mark.parametrize("key,value", [ + ("asan_runtime_dir", "/trusted/sidecar/path"), + ("hsa_asan_runtime", "/trusted/libhsa-runtime64.so"), + ("asan_extra_library_dirs", ["/trusted/llvm/lib", "/trusted/rocm_sysdeps/lib"]), +]) +def test_runtime_evidence_cannot_be_shadowed_by_config_options(tmp_path, key, value): runtime = FakeRuntime( - runtime=CapabilityCheck.ready(asan_runtime_dir="/trusted/sidecar/path") + runtime=CapabilityCheck.ready(**{key: value}) ) # Configuration parsing rejects this reserved key. Construct the immutable # object directly to retain a defense-in-depth assertion on the manager. @@ -183,7 +190,7 @@ def test_runtime_evidence_cannot_be_shadowed_by_config_options(tmp_path): base.tools[0], options={ **dict(base.tools[0].options), - "asan_runtime_dir": "/candidate/path", + key: "/candidate/path", }, ), ), @@ -193,7 +200,7 @@ def test_runtime_evidence_cannot_be_shadowed_by_config_options(tmp_path): task_config=_task(), config=config, ) - assert runtime.invocations[0][1].options["asan_runtime_dir"] == "/trusted/sidecar/path" + assert runtime.invocations[0][1].options[key] == value def test_capsule_content_is_bound_into_plan_fingerprint(tmp_path): diff --git a/tests/eval_tools/test_plugins.py b/tests/eval_tools/test_plugins.py index 38dcce559..1d1b43018 100644 --- a/tests/eval_tools/test_plugins.py +++ b/tests/eval_tools/test_plugins.py @@ -337,6 +337,40 @@ def test_gpu_asan_invocation_uses_fresh_triton_cache_and_xnack(tmp_path): assert invocation.timeout_s == 12 +@pytest.mark.parametrize("language", [KernelLanguage.TRITON, KernelLanguage.HIP]) +def test_gpu_asan_candidate_environment_matches_health(tmp_path, monkeypatch, language): + from src.eval_tools.worker import _asan_environment + + evidence = { + "host_asan_preload": "/opt/asan/llvm/lib/clang/23/lib/linux/asan.so", + "host_asan_lib_dir": "/opt/asan/llvm/lib/clang/23/lib/linux", + "hip_asan_runtime": "/opt/asan/libamdhip64.so", + "hsa_asan_runtime": "/opt/asan/libhsa-runtime64.so", + "asan_runtime_dir": "/opt/asan", + "asan_extra_library_dirs": ["/opt/asan/llvm/lib", "/opt/asan/rocm_sysdeps/lib"], + "normal_rocm_lib_dir": "/opt/rocm/lib", + } + inherited = {"LD_PRELOAD": "/existing/preload.so", "LD_LIBRARY_PATH": "/existing/lib"} + for key, value in inherited.items(): + monkeypatch.setenv(key, value) + ctx = replace( + context(tmp_path, profile(language, framework=language.value), + {"command": ["true"], **evidence}), + env=inherited, + ) + invocation = get_plugin("gpu_asan").build_invocation(ctx) + for key, value in _asan_environment(evidence).items(): + assert invocation.env[key] == value + + +@pytest.mark.parametrize("paths", ["/lib", ["relative"], [None], ["/bad\x00path"]]) +def test_gpu_asan_rejects_invalid_attested_library_paths(tmp_path, paths): + ctx = context(tmp_path, profile(KernelLanguage.HIP, framework="hip"), + {"command": ["true"], "asan_extra_library_dirs": paths}) + with pytest.raises(ValueError, match="absolute sidecar paths"): + get_plugin("gpu_asan").build_invocation(ctx) + + @pytest.mark.parametrize("tool", ["gpu_asan", "triton_fpsan", "hip_fpsan"]) def test_configured_attestation_path_is_shared_by_invocation_and_parser( tmp_path, tool diff --git a/tests/eval_tools/test_replay_capsule.py b/tests/eval_tools/test_replay_capsule.py index 9003a831f..675937b34 100644 --- a/tests/eval_tools/test_replay_capsule.py +++ b/tests/eval_tools/test_replay_capsule.py @@ -282,7 +282,8 @@ def test_flydsl_dynamic_layout_matches_no_padding_contract(): assert len(packed) == 16 # i32 + i32 + i64, with no native padding -def test_flydsl_static_launch_parser_rejects_multi_dispatch(): +@pytest.mark.parametrize("async_dependency", ["", "async [%stream] "]) +def test_flydsl_static_launch_parser_rejects_multi_dispatch(async_dependency): ir = ''' %one = arith.constant 1 : index %threads = arith.constant 128 : index @@ -290,9 +291,12 @@ def test_flydsl_static_launch_parser_rejects_multi_dispatch(): gpu.launch_func @kernels::@race blocks in (%one, %one, %one) threads in (%threads, %one, %one) dynamic_shared_memory_size %smem ''' + ir = ir.replace("gpu.launch_func @", f"gpu.launch_func {async_dependency}@") parsed = parse_flydsl_static_launch(ir) assert parsed.kernel_name == "race" assert parsed.launch.block == (128, 1, 1) assert parsed.launch.dynamic_smem_bytes == 512 with pytest.raises(CapsuleValidationError, match="requires one"): parse_flydsl_static_launch(ir + ir) + with pytest.raises(CapsuleValidationError, match="requires one"): + parse_flydsl_static_launch(ir + "\n gpu.launch_func unknown_syntax") diff --git a/tests/eval_tools/test_runtime_client.py b/tests/eval_tools/test_runtime_client.py index 451ac0c0f..a97cdadf5 100644 --- a/tests/eval_tools/test_runtime_client.py +++ b/tests/eval_tools/test_runtime_client.py @@ -108,6 +108,8 @@ def test_health_round_trip(tmp_path: Path) -> None: { "asan_runtime_dir", "hip_asan_runtime", + "hsa_asan_runtime", + "asan_extra_library_dirs", "host_asan_preload", "host_asan_lib_dir", "normal_rocm_lib_dir", diff --git a/tests/eval_tools/test_triton_aot.py b/tests/eval_tools/test_triton_aot.py new file mode 100644 index 000000000..758344358 --- /dev/null +++ b/tests/eval_tools/test_triton_aot.py @@ -0,0 +1,56 @@ +from types import SimpleNamespace + +import pytest + +from src.eval_tools.adapters.triton_aot import extract_triton_aot +from src.eval_tools.adapters.replay_capsule import CapsuleValidationError + + +def compiled(signature, constants, arg_names=None): + return SimpleNamespace( + asm={"hsaco": b"compiled-object"}, name="store", + metadata={"num_warps": 4, "warp_size": 64}, + src=SimpleNamespace(signature=signature, constants=constants, + fn=SimpleNamespace(arg_names=arg_names)), + ) + + +@pytest.mark.parametrize("signature,constants,names", [ + ({"count": "i32", "out": "*i32", "block": "i32"}, {"block": 128}, ["out", "count", "block"]), + ({"count": "i32", "out": "*i32", "block": "constexpr"}, {(2,): 128}, ["out", "count", "block"]), + ({1: "i32", 0: "*i32", 2: "i32"}, {(2,): 128}, None), + ({"1": "i32", "0": "*i32", "2": "i32"}, {"2": 128}, None), +]) +def test_extract_orders_runtime_arguments_by_source_position(tmp_path, signature, constants, names): + artifact = extract_triton_aot( + compiled(signature, constants, names), tmp_path, grid=(2,), + pointer_bindings={0: ("output", 16)}, scalar_values={1: 256}, + ) + assert [arg.name for arg in artifact.abi] == ["arg0", "arg1", "global_scratch", "profile_scratch"] + assert artifact.abi[0].ref == "output" + assert artifact.abi[0].byte_offset == 16 + assert artifact.abi[1].value == 256 + assert artifact.hsaco_path.read_bytes() == b"compiled-object" + + +@pytest.mark.parametrize("signature,constants,names,match", [ + ({"out": "*i32"}, {}, None, "no source position"), + ({"unknown": "*i32"}, {}, ["out"], "no source position"), + ({"out": "*i32"}, {}, ["out", "out"], "ambiguous"), + ({0: "*i32", "out": "*i32"}, {}, ["out"], "duplicate"), + ({0: "*i32"}, {(1, 0): 8}, None, "nested"), + ({0: "*i32"}, {"unknown": 8}, ["out"], "no source position"), +]) +def test_extract_rejects_ambiguous_or_unsupported_abi(tmp_path, signature, constants, names, match): + with pytest.raises(CapsuleValidationError, match=match): + extract_triton_aot(compiled(signature, constants, names), tmp_path, grid=(1,), + pointer_bindings={0: ("output", 0)}, scalar_values={}) + assert not list(tmp_path.iterdir()) + + +def test_extract_does_not_silently_discard_global_scratch(tmp_path): + kernel = compiled({0: "*i32"}, {}) + kernel.metadata["global_scratch_size"] = 128 + with pytest.raises(CapsuleValidationError, match="global scratch"): + extract_triton_aot(kernel, tmp_path, grid=(1,), + pointer_bindings={0: ("output", 0)}, scalar_values={}) diff --git a/tests/eval_tools/test_worker.py b/tests/eval_tools/test_worker.py index 2c45231fb..66fab809d 100644 --- a/tests/eval_tools/test_worker.py +++ b/tests/eval_tools/test_worker.py @@ -442,3 +442,79 @@ def run_step(name, argv, *, environment, **_kwargs): assert sum(value.startswith("--oracle-arg=") for value in command) == 4 assert "HSA_TOOLS_LIB" not in environment assert not any(key.startswith("RJ_CONSAN_") for key in environment) + + +@pytest.mark.parametrize( + ("failure", "expected"), + [(None, True), ("false_positive", False), ("missed_bug", False), + ("crash", False), ("uninstrumented", False)], +) +def test_triton_fpsan_requires_clean_and_bug_controls( + tmp_path: Path, monkeypatch, failure, expected +) -> None: + import json + + caches = [] + + def run_probe(name, command, *, cwd, environment, artifact_dir): + wrong = command[-1] == "wrong" + cache = Path(environment["TRITON_CACHE_DIR"]) + caches.append(cache) + cache.mkdir(parents=True) + for index in range(2): + (cache / f"kernel-{index}.json").write_text(json.dumps({ + "instrumentation_mode": "" if failure == "uninstrumented" else "fpsan" + })) + equal = not wrong + if failure == "false_positive" and not wrong: + equal = False + if failure == "missed_bug" and wrong: + equal = True + return { + "returncode": 139 if failure == "crash" else 0, + "_stdout": "AKA_FPSAN_RESULT " + json.dumps({ + "reference_digest": "reference", + "candidate_digest": "reference" if equal else "different", + }), + } + + monkeypatch.setattr(worker, "_run_probe_step", run_probe) + result = worker._triton_fpsan_positive( + tmp_path / "probes", tmp_path / "work", tmp_path / "artifacts" + ) + assert result["passed"] is expected + assert set(result["controls"]) == {"equivalent", "known_mismatch"} + assert len(set(caches)) == 2 + + +def test_gpu_asan_uses_image_owned_normal_runtime_path(tmp_path: Path, monkeypatch) -> None: + root = Path(worker.__file__).resolve().parents[2] + monkeypatch.setattr(worker, "_framework_provenance", lambda _: (root, "image", Path(worker.__file__))) + monkeypatch.setattr(worker, "_gpu_evidence", lambda: {"gpu_arch": "gfx950"}) + monkeypatch.setattr(worker, "positive_control_evidence", lambda *args, **kwargs: {"passed": False}) + monkeypatch.setenv("AKA_GPU_ASAN_NORMAL_ROCM_LIB_DIR", "/opt/rocm/lib") + runtime = tmp_path / "asan/lib" + host = runtime / "llvm/lib/clang/23/lib/linux/libclang_rt.asan-x86_64.so" + host.parent.mkdir(parents=True) + host.touch() + (runtime / "libamdhip64.so").touch() + (runtime / "libhsa-runtime64.so").touch() + (runtime / "rocm_sysdeps/lib").mkdir(parents=True) + monkeypatch.setenv("AKA_GPU_ASAN_RUNTIME_DIR", str(runtime)) + monkeypatch.setenv("LD_PRELOAD", "/custom/preload.so") + monkeypatch.setenv("LD_LIBRARY_PATH", "/custom/lib") + evidence = worker.runtime_evidence( + "gpu_asan", input_root=tmp_path, scratch_root=tmp_path, artifact_root=tmp_path + ) + assert evidence["normal_rocm_lib_dir"] == "/opt/rocm/lib" + assert "/opt/rocm/lib" in worker._asan_environment(evidence)["LD_LIBRARY_PATH"].split(":") + assert "/opt/rocm-7.2.0/lib" not in worker._asan_environment(evidence)["LD_LIBRARY_PATH"].split(":") + environment = worker._asan_environment(evidence) + assert environment["LD_PRELOAD"].split(":") == [ + str(host), str(runtime / "libamdhip64.so"), + str(runtime / "libhsa-runtime64.so"), "/custom/preload.so", + ] + assert environment["LD_LIBRARY_PATH"].split(":") == [ + str(host.parent), str(runtime), str(runtime / "llvm/lib"), + str(runtime / "rocm_sysdeps/lib"), "/opt/rocm/lib", "/custom/lib", + ] diff --git a/tests/test_docker_benchmark.sh b/tests/test_docker_benchmark.sh index f4ba375c4..0f95c4306 100755 --- a/tests/test_docker_benchmark.sh +++ b/tests/test_docker_benchmark.sh @@ -4,9 +4,10 @@ set -euo pipefail ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" RUNNER="$ROOT/src/scripts/docker_benchmark.sh" cd "$ROOT" -PINNED_GFX950_IMAGE="lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705" -PINNED_GFX950_IMMUTABLE_IMAGE="lmsysorg/sglang-rocm@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78" +PINNED_GFX950_IMAGE="lmsysorg/sglang:v0.5.20-rocm10-mi35x" +PINNED_GFX950_IMMUTABLE_IMAGE="lmsysorg/sglang@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69" export PINNED_GFX950_IMMUTABLE_IMAGE +DEFAULT_GFX950_IMAGE="lmsysorg/sglang@sha256:e20849665c105d389ef91d23c0dc73931aaa6f02056dd10e7b43e4f16c79df69" OLD_GFX950_IMAGE="lmsysorg/sglang:v0.5.12-rocm720-mi35x" REAL_PYTHON3="$(command -v python3)" export REAL_PYTHON3 @@ -180,6 +181,15 @@ fi # and the known manifest reference resolve to the same immutable local config ID, # and launch that ID rather than a mutable tag. An alias to identical content is # valid; a retagged image is rejected. +mapfile -t args < <(run_shell_args AKA_GPU_ARCH=gfx950) +assert_has "AKA_ROCM_SDK_CORE_RUNTIME=1" "${args[@]}" +for profiler_arch in gfx942 gfx1201; do + mapfile -t args < <(run_shell_args AKA_GPU_ARCH="$profiler_arch") + assert_not_has "AKA_ROCM_SDK_CORE_RUNTIME=1" "${args[@]}" +done +mapfile -t args < <(run_shell_args AKA_GPU_ARCH=gfx950 AKA_DOCKER_IMAGE=custom/image:tag) +assert_not_has "AKA_ROCM_SDK_CORE_RUNTIME=1" "${args[@]}" + mapfile -t verified_image < <( FAKE_SELECTED_IMAGE_ID=sha256:verified-config \ FAKE_PINNED_IMAGE_ID=sha256:verified-config \ @@ -200,6 +210,12 @@ mapfile -t verified_alias < <( [[ "${verified_alias[0]}" == "sha256:verified-config" ]] \ || fail "byte-identical scoring image alias was not frozen by ID" +mapfile -t alias_args < <( + AKA_GPU_ARCH=gfx950 AKA_DOCKER_IMAGE=example.invalid/scoring:alias \ + bash -c 'source "$1" smoke >/dev/null; verify_eval_tool_scoring_image; build_docker_args 0; printf "%s\n" "${docker_args[@]}"' _ "$RUNNER" +) +assert_has "AKA_ROCM_SDK_CORE_RUNTIME=1" "${alias_args[@]}" + if FAKE_SELECTED_IMAGE_ID=sha256:retagged \ FAKE_PINNED_IMAGE_ID=sha256:verified-config \ bash "$RUNNER" _verify_eval_tool_scoring_image \ @@ -207,6 +223,15 @@ if FAKE_SELECTED_IMAGE_ID=sha256:retagged \ fail "retagged scoring image unexpectedly passed immutable verification" fi +# A rollback runtime needs its matching historical sidecars; it must not be +# accepted by the ROCm 10 tool profile. +if FAKE_SELECTED_IMAGE_ID=sha256:rocm72 \ + FAKE_PINNED_IMAGE_ID=sha256:rocm10 \ + bash "$RUNNER" _verify_eval_tool_scoring_image gfx950 \ + lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705 >/dev/null 2>&1; then + fail "ROCm 7.2 scoring runtime unexpectedly passed ROCm 10 tool verification" +fi + # Artifact bind sources must be physical repository directories. Reject both a # symlink in the parent and symlinks pre-created at namespace/worker boundaries. mkdir -p "$PATH_TEST_PARENT/good" @@ -257,7 +282,8 @@ forwarded_agents="$(PATH="$FAKE_BIN:$PATH" bash "$RUNNER" _container_check_agent # The gfx950 default uses the immutable manifest, not the movable dated tag, # and retains the verified image's writable caches. mapfile -t args < <(run_shell_args AKA_GPU_ARCH=gfx950) -assert_has "$PINNED_GFX950_IMMUTABLE_IMAGE" "${args[@]}" +assert_has "$DEFAULT_GFX950_IMAGE" "${args[@]}" +assert_has "AKA_SCORING_IMAGE_REFERENCE=$DEFAULT_GFX950_IMAGE" "${args[@]}" assert_not_has "$PINNED_GFX950_IMAGE" "${args[@]}" assert_cache_args_present "" "${args[@]}" assert_not_has "AITER_ROOT_DIR=/tmp/aiter-root" "${args[@]}" @@ -272,8 +298,16 @@ mapfile -t args < <(run_shell_args AKA_GPU_ARCH=gfx950 AKA_CACHE_SUFFIX=worker/3 assert_cache_args_present "-worker_3" "${args[@]}" # Explicitly selecting the same verified tag has the same behavior. +mapfile -t args < <(run_shell_args AKA_GPU_ARCH=gfx950 AKA_SCORING_IMAGE_REFERENCE=stale-host-image) +assert_has "AKA_SCORING_IMAGE_REFERENCE=$DEFAULT_GFX950_IMAGE" "${args[@]}" +assert_not_has "AKA_SCORING_IMAGE_REFERENCE=stale-host-image" "${args[@]}" + +mapfile -t args < <(run_shell_args AKA_GPU_ARCH=gfx950 AKA_DOCKER_IMAGE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705) +assert_has "AKA_SCORING_IMAGE_REFERENCE=lmsysorg/sglang-rocm:v0.5.14-rocm720-mi35x-20260705" "${args[@]}" + mapfile -t args < <(run_shell_args AKA_GPU_ARCH=gfx950 AKA_DOCKER_IMAGE="$PINNED_GFX950_IMAGE") assert_has "$PINNED_GFX950_IMAGE" "${args[@]}" +assert_has "AKA_SCORING_IMAGE_REFERENCE=$PINNED_GFX950_IMAGE" "${args[@]}" assert_cache_args_present "" "${args[@]}" # The qualification image has the same non-root cache requirement. Test both @@ -305,7 +339,7 @@ for mode in check-agents preflight run parallel-run; do } mapfile -t args < "$TEST_HOME/deepseek-$mode.args" assert_has "$DEEPSEEK_IMAGE" "${args[@]}" - assert_not_has "$PINNED_GFX950_IMMUTABLE_IMAGE" "${args[@]}" + assert_not_has "lmsysorg/sglang-rocm@sha256:b435b508b5aa696abb25c909341ce73e41574c4271cf716bed72418dcea86b78" "${args[@]}" if [[ "$mode" != parallel-run ]]; then assert_cache_args_present "" "${args[@]}" else @@ -964,9 +998,25 @@ combined_pythonpath="$(env \ AKA_GEAK_SDK_PATH=/runtime/geak-sdk \ PYTHONPATH=/runtime/image-aiter \ bash -c "$bootstrap_script" _ python3 -c 'import os; print(os.environ["PYTHONPATH"])')" -[[ "$combined_pythonpath" == /runtime/geak-sdk:/runtime/image-aiter ]] \ +python_abi="$("$REAL_PYTHON3" -c 'import sys; print(sys.implementation.cache_tag)')" +[[ "$combined_pythonpath" == "/runtime/geak-sdk/$python_abi:/runtime/image-aiter" ]] \ || fail "GEAK bootstrap lost the image's Python import path" +# A legacy unqualified cache may contain binary wheels for the previous image. +# Only the current interpreter's cache can be imported after an upgrade. +sdk_cache="$TEST_HOME/geak-sdk-cache" +mkdir -p "$sdk_cache/$python_abi" "$sdk_cache/cpython-legacy" +printf 'value = "legacy"\n' > "$sdk_cache/aka_sdk_cache_probe.py" +printf 'value = "current"\n' > "$sdk_cache/$python_abi/aka_sdk_cache_probe.py" +printf 'value = "other-abi"\n' > "$sdk_cache/cpython-legacy/aka_sdk_cache_probe.py" +sdk_probe="$(env \ + AGENT_KERNEL_ARENA_WORKDIR="$ROOT" \ + AGENT_KERNEL_ARENA_ISOLATED_HOME=0 \ + AKA_GEAK_SDK_PATH="$sdk_cache" \ + PYTHONPATH=/runtime/image-aiter \ + bash -c "$bootstrap_script" _ "$REAL_PYTHON3" -c 'import aka_sdk_cache_probe; print(aka_sdk_cache_probe.value)')" +[[ "$sdk_probe" == current ]] || fail "GEAK imported an SDK cache for a different Python ABI" + # The explicit setup command has no run config or required agent CLI, but still # needs the GEAK-only dependency path and workflow mount for its container check. mapfile -t args < <( @@ -1193,7 +1243,7 @@ assert_has "resume" "${args[@]}" assert_has "--run-id" "${args[@]}" assert_has "run" "${args[@]}" assert_has "/sikl-config.yaml" "${args[@]}" -assert_has "AGENT_KERNEL_ARENA_IMAGE=$PINNED_GFX950_IMMUTABLE_IMAGE" "${args[@]}" +assert_has "AGENT_KERNEL_ARENA_IMAGE=$DEFAULT_GFX950_IMAGE" "${args[@]}" assert_not_has "$QUALITY_HOME/.config/gh:$QUALITY_HOME/.config/gh:ro" "${args[@]}" echo "PASS: docker_benchmark runtime, agent-selection, and eval-tool isolation tests" diff --git a/tests/test_flydsl_gfx950_vendor_compat.py b/tests/test_flydsl_gfx950_vendor_compat.py index cc289eece..0f4e552f2 100644 --- a/tests/test_flydsl_gfx950_vendor_compat.py +++ b/tests/test_flydsl_gfx950_vendor_compat.py @@ -6,13 +6,19 @@ import pytest -ROOT = Path(__file__).resolve().parents[1] / 'tasks/flydsl2flydsl' +ROOT = Path(__file__).resolve().parents[1] / 'tasks' @pytest.mark.parametrize('name', ['buffer_ops', 'vector']) @pytest.mark.parametrize('runtime', ['legacy', 'current', 'broken_dependency']) -def test_vendor_fallback_only_handles_the_removed_module(name, runtime): - source = ROOT/'flash_attn_func_kernel/flydsl_compat/__init__.py' +@pytest.mark.parametrize('task', [ + 'flydsl2flydsl/flash_attn_func_kernel', 'flydsl2flydsl/blockscale_preshuffle_gemm_kernel', + 'flydsl2flydsl/pa_decode_fp8_kernel', 'torch2flydsl/gemm_a8w8_bpreshuffle_kernel', + 'torch2flydsl/moe_sorting_kernel', 'torch2flydsl/qk_norm_rope_quant_kernel', + 'torch2flydsl/hgemm_kernel', 'torch2flydsl/jagged_dense_bmm_kernel', +]) +def test_vendor_fallback_only_handles_the_removed_module(name, runtime, task): + source = ROOT/task/'flydsl_compat/__init__.py' block = next(node for node in ast.parse(source.read_text()).body if isinstance(node, ast.Try) and isinstance(node.body[0], ast.Import) and node.body[0].names[0].name == 'flydsl.expr.' + name) diff --git a/tests/test_flydsl_task_migration_v2.py b/tests/test_flydsl_task_migration_v2.py index 6e6846b3f..8cd5e20e4 100644 --- a/tests/test_flydsl_task_migration_v2.py +++ b/tests/test_flydsl_task_migration_v2.py @@ -1473,7 +1473,7 @@ def rerun(): return result timed_run.rerun = rerun phase["value"] = "setup" - return .1, {"benchmark_method": "cuda_graph" if use_cuda_graph else "cuda_event_fallback"} + return .1, {"benchmark_method": "cuda_graph" if use_cuda_graph else "cuda_event_fallback", "benchmark_fallback_reason": fallback_reason} kmod = types.SimpleNamespace(flydsl_batched_gemm_bf16=compute, flydsl_hgemm=compute) monkeypatch.setitem(sys.modules, "aiter", types.SimpleNamespace(batched_gemm_bf16_CK=compute)) ns = {"TimedRun": Collector, "benchmark_cuda_graph_or_events": benchmark, @@ -1486,7 +1486,15 @@ def rerun(): "_retry": lambda fn, **kwargs: fn()} _harness_functions(task, {function, "_norm_worst", "_checked_gemm_output", "_gemm_reference", "_compare_gemm_output"}, ns) if behavior == "correct": - result = ns[function](verbose=False) + if name == "hgemm_kernel" and function == "arena_benchmark": + actions = module(task / "scripts/task_actions.py") + result = actions.performance(types.SimpleNamespace(arena_benchmark=ns[function])) + assert result[0]["timed_output_checked"] is True + assert result[0]["device_timing"]["benchmark_method"] == "cuda_event_fallback" + assert result[0]["device_timing"]["benchmark_fallback_reason"] == "capture_unsafe_hipblaslt_reference" + assert result[0]["device_timing"]["benchmark_method_consistent"] is True + else: + result = ns[function](verbose=False) if function == "run_benchmark":result = json.loads((tmp_path/"build/performance_report.json").read_text()) assert result[0]["timed_output_correctness"] == result[0]["replay_correctness"] == "PASS" assert calls == [(0,100,False,True),(0,100,False,False)] @@ -2271,7 +2279,8 @@ def replay(): try:return fn() finally:phase['value']='setup' timed_run.rerun=replay - return .1,{'benchmark_method':'cuda_graph' if use_cuda_graph else 'cuda_event_fallback'} + return .1,{'benchmark_method':'cuda_graph' if use_cuda_graph else 'cuda_event_fallback', + 'benchmark_fallback_reason':fallback_reason} ns={'TimedRun':Collector,'benchmark_cuda_graph_or_events':benchmark, 'require_tensor_contract':checks.require_tensor_contract,'require_unchanged':checks.require_unchanged,'verify_timed_run':checks.verify_timed_run, '_KERNEL_DIR':str(tmp_path),'KERNEL_FILE':'kernel.py','MODEL_FILE':'model.py','KERNEL_ENTRY':'flydsl_batched_gemm_a8w8', @@ -2280,7 +2289,15 @@ def replay(): 'SHAPES':[{'name':'controlled','b':1,'m':2,'n':2,'k':2}],'TOL':.01,'math':math,'json':json,'Path':Path} _harness_functions(task,{function,'_norm_worst','_checked_batched_output','_compare_batched_output'},ns) if behavior=='correct': - result=ns[function](verbose=False) + if function == 'arena_benchmark': + actions = module(task/'scripts/task_actions.py') + result = actions.performance(types.SimpleNamespace(arena_benchmark=ns[function])) + assert result[0]['timed_output_checked'] is True + assert result[0]['device_timing']['benchmark_method'] == 'cuda_event_fallback' + assert result[0]['device_timing']['benchmark_fallback_reason'] == 'capture_unsafe_aiter_hipblaslt' + assert result[0]['device_timing']['benchmark_method_consistent'] is True + else: + result=ns[function](verbose=False) if function=='run_benchmark':result=json.loads((tmp_path/'build/performance_report.json').read_text()) assert result[0]['timed_output_correctness']==result[0]['replay_correctness']=='PASS' assert calls==[(0,100,False,True),(0,100,False,False)] @@ -5669,7 +5686,9 @@ def test_torch_actual_operators_are_required_without_unused_builders(name,tmp_pa def test_unused_builder_cleanup_retains_starter_model_and_manifest_bytes(): - original={'dynamic_mxfp8_quant_kernel': {'kernel.py': '04c2ab5eb9e9bee43be84633bc7b210fcb3ad8be69bba8aa98ef6897010611a0', 'model.py': '9e5b1e289eee05aba727b71e28a98e6a7611d9fd6737d5e87b83fe9469eed39d', 'cases.json': '9e76bdcd9955b731e930c2e46536dbce2522d2f88f4ad9cb76016e825821988d'}, 'gelu_and_mul_kernel': {'kernel.py': '4fb1de9fe9d5da55e5cb924ecd03458ab70cc493612857ca300343237d541f25', 'model.py': 'c171ab0b489b1cb87a3f551c3ba8ecd820e3147a9becb6810154040f4027f7dd', 'cases.json': 'a20a152b61a241426b4f7f4d9cbdfee7af9f92f2cf1d2e8d0c2ba0aecfb20129'}, 'gelu_tanh_and_mul_kernel': {'kernel.py': '04616e2c62589d5c2e4b8147772bf4e333e753e663eb5d95bb88456428110f24', 'model.py': '95988833405bac9d10624c4ca4e78ee0251a5a60c1dd457d6b915904b9b2dadf', 'cases.json': 'a20a152b61a241426b4f7f4d9cbdfee7af9f92f2cf1d2e8d0c2ba0aecfb20129'}, 'gemm_a8w8_bpreshuffle_kernel': {'kernel.py': 'b5d3e87a3ceca3fe555572f0b4ab5c7b1dd6c3f5c9e18ab924b589df3d485996', 'model.py': 'e4278a3637b56eab94baec3b712ff1fb5ac44206a0ac21b1925a7497ca0b9643', 'cases.json': '7f2d3e25a614486974800da54cb23d2c7913101b2ec1b82ca92818cdfec479ee'}, 'hgemm_kernel': {'kernel.py': '29cab2057d32224a7da558083f6b4edeb560e44efd59a60b82dfb14c9c6a8d28', 'model.py': '89ca4fce55817fdf5fcbaea925a1639f9b96cd809dcda8c06cd70f9e1033372e', 'cases.json': '8388fcafea635e69bde93aad82d8b6bcd10d3ffd9998f2594e4ab989c2fd61e4'}, 'jagged_dense_bmm_kernel': {'kernel.py': 'fcb9b75ec238ced56568fb5b27535a160314db29c212abe33c83be6f7df3c043', 'model.py': '1b446ea35fee03f47ae16121fb7ba8aa933e9e99c48f2d90584c770d56065186', 'cases.json': '8304b063316f9cfc3667a8f38d9d85805b38a2339deebd48b883bfb59b5427e3'}, 'moe_sorting_kernel': {'kernel.py': '4bc536d6d29f16e1f24278d9db05ffba13d723f3f42064cce18b60c687f3eb43', 'model.py': 'c874911efc1c947437d5c7c62019e7c52b34458a0ec1560ecdde72aa971de271', 'cases.json': '0c8247d8727d48dc3a8ddd20ede1db3ae2586e990222f7bf19ce39dc6ad913c4'}, 'qk_norm_rope_quant_kernel': {'kernel.py': 'be6c11328764ac59054301e1d8a9312287227718eee498731698ddc45079fa46', 'model.py': '152a32140302f1264c555fbd9b6f9d8362583fd08b0344290a67d6b1bb849ee2', 'cases.json': 'de8e183a711424ddabe5f8bfea4fa00dbe79cb886462b5d3681f67fa9d6e0eb1'}, 'rmsnorm2d_dynamicquant_kernel': {'kernel.py': 'c8542e7ea4b69aa6881ae2e7046995c67112e7e38b68e95d3cc2a51965d05bdc', 'model.py': '5cd3abaf088651f8f2ed9db44513bfd8c5238a45145aaf3dcd3168d63e3ff080', 'cases.json': '046dcf1c5e6f49bf68bd6e935555006803f9f6af5668460389ae6147297528f1'}, 'rmsnorm2d_kernel': {'kernel.py': 'a20840e12a22f5c08fed1e87eee62de5dec590680c79b5b6fa8cb28d9dd9396a', 'model.py': '8442cb4d63444e7dfa9db1fa2d6253ceee0463b1debfd6919fc5219181de3b13', 'cases.json': '426c9e84d97161e5bb7a09353102c57648f5ca5a4790c46370b5369ba470fa64'}, 'rmsnorm2d_smoothquant_kernel': {'kernel.py': '701e11ec5e63bf65572fd9325a7e0b1bea0e1aa9561f8711de0e7ed0113fc2b0', 'model.py': '5496289dc3f8f72c1b9deefce1e39b7c1a00dc5726ad708abe65641fee4edf0e', 'cases.json': '359795558e6ebcfb617bbae66eda8540f5503cc5a04056f1fdfde58e13aedce4'}, 'swiglu_and_mul_kernel': {'kernel.py': '6adbabe7f43dd51289ec4afa3310ba30c56ad8f216edb801886882a1c00715e5', 'model.py': '74765a3a6e469d27926a92b0d710231f2bb7604850186fca782dc5446a8224b4', 'cases.json': 'a20a152b61a241426b4f7f4d9cbdfee7af9f92f2cf1d2e8d0c2ba0aecfb20129'}} + # Starter hashes include the reviewed FlyDSL 0.3.2 API routing changes. + # Model and workload manifest hashes remain those of the original tasks. + original={'dynamic_mxfp8_quant_kernel': {'kernel.py': '04c2ab5eb9e9bee43be84633bc7b210fcb3ad8be69bba8aa98ef6897010611a0', 'model.py': '9e5b1e289eee05aba727b71e28a98e6a7611d9fd6737d5e87b83fe9469eed39d', 'cases.json': '9e76bdcd9955b731e930c2e46536dbce2522d2f88f4ad9cb76016e825821988d'}, 'gelu_and_mul_kernel': {'kernel.py': '4fb1de9fe9d5da55e5cb924ecd03458ab70cc493612857ca300343237d541f25', 'model.py': 'c171ab0b489b1cb87a3f551c3ba8ecd820e3147a9becb6810154040f4027f7dd', 'cases.json': 'a20a152b61a241426b4f7f4d9cbdfee7af9f92f2cf1d2e8d0c2ba0aecfb20129'}, 'gelu_tanh_and_mul_kernel': {'kernel.py': '04616e2c62589d5c2e4b8147772bf4e333e753e663eb5d95bb88456428110f24', 'model.py': '95988833405bac9d10624c4ca4e78ee0251a5a60c1dd457d6b915904b9b2dadf', 'cases.json': 'a20a152b61a241426b4f7f4d9cbdfee7af9f92f2cf1d2e8d0c2ba0aecfb20129'}, 'gemm_a8w8_bpreshuffle_kernel': {'kernel.py': 'da4cc1b465b8f5c620c88e5876a2e4799023b5d585a00b2f203e3914c2379269', 'model.py': 'e4278a3637b56eab94baec3b712ff1fb5ac44206a0ac21b1925a7497ca0b9643', 'cases.json': '7f2d3e25a614486974800da54cb23d2c7913101b2ec1b82ca92818cdfec479ee'}, 'hgemm_kernel': {'kernel.py': '27d4c78ecd9c8e42507f8c801ea80ab6e0fbd7667d0672e02b85a0b8ebe34667', 'model.py': '89ca4fce55817fdf5fcbaea925a1639f9b96cd809dcda8c06cd70f9e1033372e', 'cases.json': '8388fcafea635e69bde93aad82d8b6bcd10d3ffd9998f2594e4ab989c2fd61e4'}, 'jagged_dense_bmm_kernel': {'kernel.py': 'dbbda2a72ebe1872cf2a286eea086a33702c0bc4646a8e2006edc37be5fa7f0a', 'model.py': '1b446ea35fee03f47ae16121fb7ba8aa933e9e99c48f2d90584c770d56065186', 'cases.json': '8304b063316f9cfc3667a8f38d9d85805b38a2339deebd48b883bfb59b5427e3'}, 'moe_sorting_kernel': {'kernel.py': 'e46cc487791f85ae3aab7d32955f2b0784c946f1d3f7d8a03b89de53a7d3c8e0', 'model.py': 'c874911efc1c947437d5c7c62019e7c52b34458a0ec1560ecdde72aa971de271', 'cases.json': '0c8247d8727d48dc3a8ddd20ede1db3ae2586e990222f7bf19ce39dc6ad913c4'}, 'qk_norm_rope_quant_kernel': {'kernel.py': '1ca597964eb0b3899435b60fc01509891613a13cbe0d3552c6feac803ffe4d2a', 'model.py': '152a32140302f1264c555fbd9b6f9d8362583fd08b0344290a67d6b1bb849ee2', 'cases.json': 'de8e183a711424ddabe5f8bfea4fa00dbe79cb886462b5d3681f67fa9d6e0eb1'}, 'rmsnorm2d_dynamicquant_kernel': {'kernel.py': 'c8542e7ea4b69aa6881ae2e7046995c67112e7e38b68e95d3cc2a51965d05bdc', 'model.py': '5cd3abaf088651f8f2ed9db44513bfd8c5238a45145aaf3dcd3168d63e3ff080', 'cases.json': '046dcf1c5e6f49bf68bd6e935555006803f9f6af5668460389ae6147297528f1'}, 'rmsnorm2d_kernel': {'kernel.py': 'a20840e12a22f5c08fed1e87eee62de5dec590680c79b5b6fa8cb28d9dd9396a', 'model.py': '8442cb4d63444e7dfa9db1fa2d6253ceee0463b1debfd6919fc5219181de3b13', 'cases.json': '426c9e84d97161e5bb7a09353102c57648f5ca5a4790c46370b5369ba470fa64'}, 'rmsnorm2d_smoothquant_kernel': {'kernel.py': '701e11ec5e63bf65572fd9325a7e0b1bea0e1aa9561f8711de0e7ed0113fc2b0', 'model.py': '5496289dc3f8f72c1b9deefce1e39b7c1a00dc5726ad708abe65641fee4edf0e', 'cases.json': '359795558e6ebcfb617bbae66eda8540f5503cc5a04056f1fdfde58e13aedce4'}, 'swiglu_and_mul_kernel': {'kernel.py': '6adbabe7f43dd51289ec4afa3310ba30c56ad8f216edb801886882a1c00715e5', 'model.py': '74765a3a6e469d27926a92b0d710231f2bb7604850186fca782dc5446a8224b4', 'cases.json': 'a20a152b61a241426b4f7f4d9cbdfee7af9f92f2cf1d2e8d0c2ba0aecfb20129'}} for name,files in original.items(): for rel,expected in files.items(): assert hashlib.sha256((ROOT/'tasks/torch2flydsl'/name/rel).read_bytes()).hexdigest()==expected,(name,rel) @@ -6669,7 +6688,7 @@ def rerun(): timed_run.rerun=rerun else:fn() phase['name']='setup' - return .1,{'benchmark_method':'cuda_event_fallback','benchmark_timed_run_kind':'eager_callable'} + return .1,{'benchmark_method':'cuda_event_fallback','benchmark_timed_run_kind':'eager_callable','benchmark_fallback_reason':fallback_reason} monkeypatch.setattr(torch.cuda,'synchronize',lambda:None);monkeypatch.setattr(torch.cuda,'empty_cache',lambda:None) ns={'TimedRun':Collector,'benchmark_cuda_graph_or_events':bench,'require_tensor_contract':checks.require_tensor_contract,'require_unchanged':checks.require_unchanged,'verify_timed_run':checks.verify_timed_run,'math':math,'json':json,'Path':Path, '_KERNEL_DIR':str(tmp_path),'KERNEL_FILE':'kernel.py','MODEL_FILE':'model.py','_make_inputs':lambda *a:(x,w),'_load_module':lambda directory,filename,alias:mmod if filename=='model.py' else kmod, @@ -6678,7 +6697,15 @@ def rerun(): # Measured-only controls are separately tested on both real timing paths. should_pass=behavior=='correct' or correctness and behavior in {'measured_wrong','replay_wrong','cached','wrong_scale'} if should_pass: - result=ns[function](verbose=False) + if function == 'arena_benchmark': + actions = module(task/'scripts/task_actions.py') + result = actions.performance(types.SimpleNamespace(arena_benchmark=ns[function])) + assert result[0]['timed_output_checked'] is True + assert result[0]['device_timing']['benchmark_method'] == 'cuda_event_fallback' + assert result[0]['device_timing']['benchmark_fallback_reason'] == 'capture_unsafe_hipblaslt_reference' + assert result[0]['device_timing']['benchmark_method_consistent'] is True + else: + result=ns[function](verbose=False) if not correctness: report=json.loads((tmp_path/'build/performance_report.json').read_text()) if function=='run_benchmark' else result assert report[0]['timed_output_correctness']==report[0]['replay_correctness']=='PASS' @@ -7159,7 +7186,7 @@ def test_qk_original_source_passes_static_compile_without_kernel_or_case_edit(): assert result.metadata['compile_kind']=='python_bytecode' # Syntax evidence only: GPU compilation/correctness/performance are required # again with the corrected dependency policy and unchanged old API source. - assert hashlib.sha256((task/'kernel.py').read_bytes()).hexdigest()=='be6c11328764ac59054301e1d8a9312287227718eee498731698ddc45079fa46' + assert hashlib.sha256((task/'kernel.py').read_bytes()).hexdigest()=='1ca597964eb0b3899435b60fc01509891613a13cbe0d3552c6feac803ffe4d2a' @@ -7294,3 +7321,63 @@ def test_a8w8_production_baseline_uses_one_consistent_reduction_policy(): assert seen=={'run_correctness':1,'run_benchmark':1,'arena_benchmark':1} cfg=load_task_spec(task/'config.yaml',task_id='torch2flydsl/gemm_a8w8_kernel') assert cfg.baseline.correctness_policy=='required' + + +@pytest.mark.parametrize('source', [ + 'import sys\nsys.modules["arena_harness"]._norm_worst = lambda *a: (0, 0)', + 'import sys as runtime\nruntime.modules["arena_harness"].TOL = 1e9', + 'from sys import modules as loaded\nloaded["arena_harness"]._compare_batched_output = lambda *a: None', + 'import sys\ngetattr(sys, "modules")["arena_harness"].TOL = 1e9', + 'import sys\nvars(sys)["modules"]["arena_harness"].TOL = 1e9', + 'import sys\nsys.__dict__["modules"]["arena_harness"].TOL = 1e9', + 'import inspect\ninspect.currentframe().f_back.f_globals["TOL"] = 1e9', + 'def helper(): pass\nhelper.__globals__["TOL"] = 1e9', + 'from builtins import getattr as lookup\nlookup(obj, "modules")', + 'from sys import setprofile as disable\ndisable(None)', +]) +def test_batched_int8_rejects_comparator_state_bypass_before_import(tmp_path, source): + task = ROOT / 'tasks/torch2flydsl/batched_gemm_a8w8_kernel' + runtime = module(task / 'task_runtime.py') + candidate = tmp_path / 'kernel.py' + candidate.write_text('import flydsl\n' + source + '\n') + with pytest.raises(ValueError, match='Protected|introspection'): + runtime.check_dependencies([candidate], final_language=True) + + +def test_batched_int8_allows_ordinary_kernel_preparation_dependencies(tmp_path): + runtime = module(ROOT / 'tasks/torch2flydsl/batched_gemm_a8w8_kernel/task_runtime.py') + candidate = tmp_path / 'kernel.py' + candidate.write_text('import flydsl.compiler as flyc\nimport torch\nimport math\nimport sys\nlimit = sys.maxsize\n') + runtime.check_dependencies([candidate], final_language=True) + + +@pytest.mark.parametrize('source', [ + 'import torch\ntorch.matmul = lambda *a: 0', + 'import torch as t\nt.matmul = lambda *a: 0', + 'import torch\nt = torch\nu = t\nu.matmul = lambda *a: 0', + 'from torch import testing as checks\nchecks.assert_close = lambda *a: None', + 'import torch\ndel torch.matmul', + 'import torch\ntorch.__dict__["matmul"] = lambda *a: 0', + 'import torch\nsetattr(torch, "matmul", lambda *a: 0)', + 'import torch\nwrite = setattr\nwrite(torch, "matmul", lambda *a: 0)', + 'import torch\ngetattr(torch, "__dict__")["matmul"] = lambda *a: 0', + 'import sys\nsys.modules["torch"].matmul = lambda *a: 0', + 'from sys import modules as loaded\nloaded["torch"].matmul = lambda *a: 0', + 'import inspect\ninspect.currentframe().f_back.f_globals["TOL"] = 1e9', + 'def helper(): pass\nhelper.__globals__["TOL"] = 1e9', +]) +def test_hgemm_rejects_oracle_mutation_before_import(tmp_path, source): + runtime = module(ROOT / 'tasks/torch2flydsl/hgemm_kernel/task_runtime.py') + candidate = tmp_path / 'kernel.py' + candidate.write_text('import flydsl\n' + source + '\n') + with pytest.raises(ValueError, match='Protected|introspection'): + runtime.check_dependencies([candidate], final_language=True) + + +def test_hgemm_retains_local_compilation_cache_and_initial_kernel(tmp_path): + task = ROOT / 'tasks/torch2flydsl/hgemm_kernel' + runtime = module(task / 'task_runtime.py') + runtime.check_dependencies([task / 'kernel.py'], final_language=True) + candidate = tmp_path / 'kernel.py' + candidate.write_text('import flydsl\ndef run(exe):\n cf = getattr(exe, "_cf", None)\n exe._cf = cf\n') + runtime.check_dependencies([candidate], final_language=True) diff --git a/tests/test_mla_runtime_compat.py b/tests/test_mla_runtime_compat.py new file mode 100644 index 000000000..cae31dbc4 --- /dev/null +++ b/tests/test_mla_runtime_compat.py @@ -0,0 +1,40 @@ +"""Exercise the gfx950 MLA pipeline regression on actual HIP hardware.""" +import shutil +import subprocess +import sys +from pathlib import Path + +import pytest + +from src.perf_helper_materialization import materialize_perf_helpers_in_workspace + + +def test_gfx950_mla_repeated_decode_matches_reference(tmp_path): + torch = pytest.importorskip("torch") + if not torch.version.hip or not torch.cuda.is_available(): + pytest.skip("requires a ROCm GPU") + if not torch.cuda.get_device_properties(0).gcnArchName.startswith("gfx950"): + pytest.skip("requires gfx950") + source = Path(__file__).resolve().parents[1] / "tasks/triton2flydsl/aiter/mla" + task = tmp_path / "mla" + shutil.copytree(source, task) + materialize_perf_helpers_in_workspace(task) + # The harness changes cwd and owns task-local imports; isolate its module state. + subprocess.run([sys.executable, "-c", """ +import torch +import test_kernel_harness as h +m = h.load_module() +shape = h.TEST_SHAPES[3] +for seed in (7, 45, 46, 123): + torch.manual_seed(seed) + q, kv, out, table, cu, lengths, scale = h.make_test_data(*shape) + reference = h.torch_mla_extend(q, kv, cu, lengths, table, shape[4], scale, + o_dtype=q.dtype) + for repeat in range(3): + out.fill_(float('nan')) + actual = h._call_kernel(m, q, kv, out, cu, lengths, shape[-1], table, + scale, shape[4], shape[5]) + torch.cuda.synchronize() + h._checked_mla_output(actual, out, q, shape[4]) + h._compare_mla_output(actual, reference) +"""], cwd=task, check=True, timeout=180) diff --git a/tests/test_rocm_sdk_runtime.py b/tests/test_rocm_sdk_runtime.py new file mode 100644 index 000000000..416120b96 --- /dev/null +++ b/tests/test_rocm_sdk_runtime.py @@ -0,0 +1,46 @@ +from types import SimpleNamespace + +import pytest + +from src.scripts import rocm_sdk_runtime as wrapper + + +def test_profiler_uses_core_tree_without_losing_workload_arguments(tmp_path, monkeypatch): + core = tmp_path / "_rocm_sdk_core" + executable = core / "bin" / "rocprofv3" + executable.parent.mkdir(parents=True) + executable.touch() + monkeypatch.setattr(wrapper.importlib.util, "find_spec", + lambda name: SimpleNamespace(origin=str(core / "__init__.py"))) + devel = tmp_path / "_rocm_sdk_devel" / "lib" + env = {"LD_LIBRARY_PATH": f"{devel}:/custom/lib:{core / 'lib'}", "ROCR_VISIBLE_DEVICES": "2"} + arguments = ["--kernel-trace", "--", "python", "workload with spaces.py"] + + command, actual_env = wrapper.profiler_command(arguments, env) + + assert command == [str(executable), *arguments] + assert actual_env == {**env, "LD_LIBRARY_PATH": f"{core / 'lib'}:/custom/lib"} + assert env["LD_LIBRARY_PATH"].startswith(str(devel)) + + +def test_missing_core_profiler_does_not_fall_back_to_conflicting_devel(tmp_path, monkeypatch): + monkeypatch.setattr(wrapper.importlib.util, "find_spec", + lambda name: SimpleNamespace(origin=str(tmp_path / "__init__.py"))) + with pytest.raises(RuntimeError, match="profiler is missing"): + wrapper.profiler_command([], {}) + + +def test_workload_environment_does_not_require_a_profiler(tmp_path, monkeypatch): + core = tmp_path / "_rocm_sdk_core" + monkeypatch.setattr(wrapper.importlib.util, "find_spec", + lambda name: SimpleNamespace(origin=str(core / "__init__.py"))) + original = {"LD_LIBRARY_PATH": "/custom", "ROCR_VISIBLE_DEVICES": "3"} + assert wrapper.runtime_environment(original) == { + **original, "LD_LIBRARY_PATH": f"{core / 'lib'}:/custom", + } + + +def test_missing_core_fails_closed(monkeypatch): + monkeypatch.setattr(wrapper.importlib.util, "find_spec", lambda name: None) + with pytest.raises(RuntimeError, match="no ROCm SDK core"): + wrapper.runtime_environment({}) diff --git a/tests/test_task_materialization_v2.py b/tests/test_task_materialization_v2.py index bccdfa514..ea085f204 100644 --- a/tests/test_task_materialization_v2.py +++ b/tests/test_task_materialization_v2.py @@ -399,12 +399,27 @@ def test_setup_cannot_change_contract_or_produce_completion_evidence(tmp_path, c assert record_unchecked(expected_workspace(tmp_path))["status"] == "failed" -def test_resume_checks_runtime_identity(tmp_path, monkeypatch): - monkeypatch.setenv("AKA_SCORING_IMAGE_RUNTIME_REF", "runtime@sha256:original") +@pytest.mark.parametrize("identity_variable", [ + "AKA_SCORING_IMAGE_RUNTIME_REF", "AKA_SCORING_IMAGE_REFERENCE", +]) +def test_resume_checks_runtime_identity(tmp_path, monkeypatch, identity_variable): + monkeypatch.setenv(identity_variable, "runtime@sha256:original") config, _ = task(tmp_path) workspace = run(config, tmp_path) write(workspace, "source/kernel.py", "optimized") - monkeypatch.setenv("AKA_SCORING_IMAGE_RUNTIME_REF", "runtime@sha256:changed") + monkeypatch.setenv(identity_variable, "runtime@sha256:changed") + with pytest.raises(MaterializationError, match="runtime identity changed"): + run(config, tmp_path) + assert (workspace / "source/kernel.py").read_text() == "optimized" + + +def test_resume_rejects_unbound_legacy_runtime_without_touching_candidate(tmp_path, monkeypatch): + monkeypatch.delenv("AKA_SCORING_IMAGE_RUNTIME_REF", raising=False) + monkeypatch.delenv("AKA_SCORING_IMAGE_REFERENCE", raising=False) + config, _ = task(tmp_path) + workspace = run(config, tmp_path) + write(workspace, "source/kernel.py", "optimized") + monkeypatch.setenv("AKA_SCORING_IMAGE_REFERENCE", "runtime@sha256:bound") with pytest.raises(MaterializationError, match="runtime identity changed"): run(config, tmp_path) assert (workspace / "source/kernel.py").read_text() == "optimized"