diff --git a/CHANGELOG.md b/CHANGELOG.md index 38fcd8db47..a61800baac 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -2,24 +2,26 @@ # CUTLASS 4.x -## [4.8.0](https://github.com/NVIDIA/cutlass) (2026-08-25) +## [4.8.0](https://github.com/NVIDIA/cutlass/releases/tag/v4.8.0) (2026-09-17) ### CuTe DSL * New features - Initial Rubin support to accelerate dense GEMMs. The following features are available: - - CuTe DSL - - Rubin new FP8 and FP4 Tensor Core support + - CuTe DSL and CuTe extensions + - Support for higher-throughput FP8 (MMA_K=64) and FP4 (MMA_K=128) Tensor Core MMA instructions - B collector reuse - Extended TMEM size from 512 COL to 576 COL - Larger shared memory allocations (328KB) - Enhanced mixed precision throughput (FP8/FP4) + - Softmax acceleration related features - Primitives - - Rubin new FP8 and FP4 Tensor Core support + - Support for higher-throughput FP8 (MMA_K=64) and FP4 (MMA_K=128) Tensor Core MMA instructions + - B collector reuse - Extended TMEM size from 512 COL to 576 COL - - NOTE: Executing Rubin kernels (SM107) requires the R615 driver which will be released - with CUDA Toolkit 13.4 GA. R610 from CUDA Toolkit 13.4 Developer Preview is not - sufficient. + - Larger shared memory allocations (328KB) + - Enhanced mixed precision throughput (FP8/FP4) + - Softmax acceleration related feature + - 2:4 sparsity support for FP4 - CuTe DSL extensions has several new features: - CTA-V maps are now inferred automatically for `cute_ext` TMA load, store, multicast, and reduce-store operations. Explicit CTA-V maps remain supported as overrides. @@ -29,7 +31,7 @@ - Improved device-side TMA descriptor updates and grouped GEMM performance through SMEM-staged updates, workspace reuse, and reduced prologue and synchronization overhead. - This release includes an opt-in preview of the CuTe DSL extensions (`cute_ext`) compiler pipeline for ordinary Cute DSL kernels. This pipeline lets user mix `cute_ext` APIs directly into `@cute.jit` and `@cute.kernel` code and is required for kernels that mix the two API surfaces. You may test this feature with the following: `CUTE_DSL_USE_EXTENSION_COMPILER=1 python your_program.py` -The pipeline is expected to preserve program behavior, but generated PTX/SASS may differ. The pipeline is planned to become the default in a future release. +The pipeline is expected to preserve program behavior and performance, but generated PTX/SASS may differ. Note that this pipeline will become the default in the future, no earlier than 4.10. - Added examples for better control over Primitives' compiler warnings/errors introduced in 4.7.0. See the `CuTeDSL/experimental/compiler_diagnostic/` directory. - IKET Profiler Tool - Rubin kernels (sm107) can now be profiled. @@ -43,7 +45,7 @@ The pipeline is expected to preserve program behavior, but generated PTX/SASS ma - Grouped blockscaled GEMM with B collector reuse as applicable - Blockwise GEMM - Rubin (CuTe extension): - - FP4 blockscaled GEMM + - Support for higher-throughput FP8 (MMA_K=64) and FP4 (MMA_K=128) blockscaled GEMM with UE5M3 scale-factor - Grouped GEMM with B collector reuse - Blackwell (CuTe extension): - Dense GEMMs @@ -57,6 +59,11 @@ The pipeline is expected to preserve program behavior, but generated PTX/SASS ma - Input transform GEMM - GeForce pingpong dense GEMM - Blackwell Ultra blockscaled GEMM + - Dense Convolutions + - Implicit-Gemm Fprop Conv + - Blocksclaed Implicit-Gemm Fprop Conv + - GeForce Implicit-Gemm Fprop Conv + - GeForce Blockscaled Implicit-Gemm Fporp Conv - Attention - GQA Decode - Grouped GEMM @@ -64,6 +71,11 @@ The pipeline is expected to preserve program behavior, but generated PTX/SASS ma - Top-K - Ampere (CuTe extension): - SIMT GEMM + - CuTe DSL now supports x86_64 Windows + - CuTe DSL AoT now supports new host target: QNX8.0 + - Notebooks are restructured under examples/python/CuTeDSL/cute/notebooks and new notebooks for primitives will be added under examples/python/CuTeDSL/notebooks + - Numpy is now not a default dependency + * Bug fixes and improvements: - `nvidia-cuda-nvdisasm` is now an optional dependency of `nvidia-cutlass-dsl` via the optional `[sass]` extra. SASS dumping (`CUTE_DSL_KEEP=sass` / KeepSASS) now resolves `nvdisasm` from the bundled wheel (recommended since its version matches the DSL toolchain) or from a local CUDA Toolkit (`CUDA_HOME`/`CUDA_PATH`). A locally-provided `nvdisasm` must come from a CUDA Toolkit at least as new as the toolchain that produced the CUBIN. Installations that never dump SASS are unaffected. - Reduced the protobuf version requirement of IKET profiler from 6.30 to 4.21. This should make protobuf an easier requirement to satisfy in preparation for transioning IKET to an optional extra. @@ -72,49 +84,61 @@ The pipeline is expected to preserve program behavior, but generated PTX/SASS ma vectorized instructions for tensors with a dynamic stride ([!3463](https://github.com/NVIDIA/cutlass/issues/3463)) - Fixed TVM-FFI env stream detection for GPU tensors in tuple ([!3444](https://github.com/NVIDIA/cutlass/issues/3444)) + - Fixed GPU `link-libraries` compile-option order so it is stable across processes + ([!3564](https://github.com/NVIDIA/cutlass/issues/3564)) + - Fixed preprocessor `IndexError` on staged `bool()` with no arguments + ([!3506](https://github.com/NVIDIA/cutlass/issues/3506)) + - Rejected `cute.compile` on `@cute.kernel` with a user error instead of an ICE + ([!3429](https://github.com/NVIDIA/cutlass/issues/3429)) + - Fixed CuTe DSL crashing the Python interpreter when used in a REPL + ([!3413](https://github.com/NVIDIA/cutlass/issues/3413)) + - Fixed a cuDNN Frontend FROST SDPA backward compilation failure issue ([!3594](https://github.com/NVIDIA/cutlass/issues/3594)) + - Fixed SIGABRTs in TVM-FFI launch for cuDNN Frontend SM100 ragged SDPA kernel ([!3595](https://github.com/NVIDIA/cutlass/issues/3595)) This release has been tested against the following packages: - - FlashAttention: [main (0251105)](https://github.com/Dao-AILab/flash-attention/commit/0251105a2fb19d2957484b7f023cd8c115286ced) - - Quack: [main (60d8808)](https://github.com/Dao-AILab/quack/commit/60d88082272a256fa9b3b2ab631c82cfa78337c6) - - FlashInfer: [main (109d44f)](https://github.com/flashinfer-ai/flashinfer/commit/109d44fceea027290d54efcfe927f8a5665b59de) - - cuDNN-Frontend: [deveop (25b3d51)](https://github.com/NVIDIA/cudnn-frontend/commit/25b3d5126b6544afc209e3c2e94a74f5d82db201) - - Pytorch: [main (cf30153)](https://github.com/pytorch/pytorch/commit/cf30153c4c131c8164ee7798e5022d810682e2cb) - - TensorRT-LLM: [main (1cef02e)](https://github.com/NVIDIA/TensorRT-LLM/commit/1cef02e901be43081b1ba6d4981e94ed3bd9c1e8) + - FlashAttention: [main (8d3a3b8)](https://github.com/Dao-AILab/flash-attention/commit/8d3a3b80d4758ebde5a867c50d24d4351443cf2b) + - Quack: [main (35266c3)](https://github.com/Dao-AILab/quack/commit/35266c3298f0e9bf6d5f46c30aace2eaeae517e3) + - FlashInfer: [main (5d0c89e)](https://github.com/flashinfer-ai/flashinfer/commit/5d0c89eacae6ca08f2a1ce92eba557bbad7a1bfc) + - cuDNN-Frontend: [deveop (e0317d1)](https://github.com/NVIDIA/cudnn-frontend/commit/e0317d1f6cbc1bb8e50abce5870b6bf59b9d4a74) + - Pytorch: [main (7d5f021)](https://github.com/pytorch/pytorch/commit/7d5f0216450b8253e62d88284d49de725d88bd42) + - TensorRT-LLM: [main (c295dd9)](https://github.com/NVIDIA/TensorRT-LLM/commit/c295dd9fca143a3fdd2769617699fdd4d4a93eb5) ### CUTLASS Operator API -* Operator API features and functionality: - - Dense and blockscaled GEMMs in Operator API have preliminary Rubin support. These are provided as a preview and may need additional performance tuning. +* Dense and blockscaled GEMMs in Operator API have preliminary Rubin support. These are provided as a preview and may need additional performance tuning. + + - Updated GEMMs include: + + - Dense GEMMs: FP8xFP8 + - Blockscaled GEMM: {MXFP8}x{MXFP4, MXFP8} and {MXFP4, NVFP4}x{MXFP4, NVFP4} (including support for the new UE5M3 scale factor dtype for NVFP4). + + + - These kernels utilize the below new features in Rubin: + + - Higher SMEM (328KB) and TMEM capacity (288KB) + - B-buffer reuse + - Enhanced mixed precision throughput - Updated GEMMs include: - - Dense GEMMs: FP8xFP8 - - Blockscaled GEMM: {MXFP8}x{MXFP4, MXFP8} and {MXFP4, NVFP4}x{MXFP4, NVFP4} (including support for the new UE5M3 scale factor dtype for NVFP4). - - These kernels utilize the below new features in Rubin: - - Higher SMEM (328KB) and TMEM capacity (288KB) - - B-buffer reuse - - Enhanced mixed precision throughput +* Operators can now be ranked by their estimated performance when nvMatmulHeuristics is available. See tutorial here. NOTE: This currently only supports Blackwell kernels as nvMatmulHeuristics does not yet support Rubin. - - Operators can now be ranked by their estimated performance when nvMatmulHeuristics is available. See tutorial [here](https://docs.nvidia.com/cutlass/latest/media/docs/operators/tutorials/007_heuristics.ipynb) +* Standalone kernel implementations are now exposed through `cutlass.kernels`. These kernels can be used directly, in addition to being discoverable and usable via the Operator interface in `cutlass.operators`. - NOTE: This currently only supports Blackwell kernels as nvMatmulHeuristics does not yet support Rubin. +* Custom Epilogue fusions now support per-row or per-column reductions. - - Standalone kernel implementations are now exposed through `cutlass.kernels`, in addition to those exposed through the Operator interface in `cutlass.operators`. This allows kernels to be called directly without looking them up first. - - Custom Epilogue fusions now support partial (per-row or per-column) reductions. - - `IndexPtrGroupedGemmArguments` is now used to represent Grouped GEMM with contiguous-offset/index-pointers. Existing `GroupedGemmArguments` is deprecated and will be removed in a future release. +* `IndexPtrGroupedGemmArguments` is now used to represent Grouped GEMM with contiguous-offset/index-pointers. Existing `GroupedGemmArguments` is deprecated and will be removed in a future release. ### C++ * Added initial Rubin support (SM107) with CuTe C++ building blocks: - - [Rubin Tensor Core MMA instructions](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/mma_sm107_umma.hpp) and corresponding [CuTe MMA traits](https://github.com/NVIDIA/cutlass/blob/main/include/cute/atom/mma_traits_sm107.hpp). + - [Rubin Tensor Core MMA instructions](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/mma_sm107_umma.hpp) and corresponding [CuTe MMA traits](https://github.com/NVIDIA/cutlass/blob/main/include/cute/atom/mma_traits_sm107.hpp). * CuTe examples that demonstrate the use of Rubin SM107 Tensor Core instructions: - - [Dense FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8.cu). - - [Block-scaled FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8_blockscaled.cu). - - [Mixed-precision block-scaled FP8/FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8_blockscaled.cu). - - [Block-scaled FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp4_blockscaled.cu). + - [Dense FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8.cu). + - [Block-scaled FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8_blockscaled.cu). + - [Mixed-precision block-scaled FP8/FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8_blockscaled.cu). + - [Block-scaled FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp4_blockscaled.cu). * Adjusted shared-memory and tensor-memory capacity handling for Rubin SM107: - - Set the [SM107 shared-memory capacity](https://github.com/NVIDIA/cutlass/blob/main/include/cutlass/arch/arch.h) to 327 KiB and added launch support for oversized shared-memory configurations. - - Set the [SM107 TMEM capacity](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/tmem_capacity_sm100.hpp) to 576 columns per SM, updated the [CuTe 1SM and 2SM TMEM allocators](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/tmem_allocator_sm100.hpp) for Rubin's exclusive allocation path. + - Set the [SM107 shared-memory capacity](https://github.com/NVIDIA/cutlass/blob/main/include/cutlass/arch/arch.h) to 327 KiB and added launch support for oversized shared-memory configurations. + - Set the [SM107 TMEM capacity](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/tmem_capacity_sm100.hpp) to 576 columns per SM, updated the [CuTe 1SM and 2SM TMEM allocators](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/tmem_allocator_sm100.hpp) for Rubin's exclusive allocation path. * Enabled the existing SM100-compatible [GEMM](https://github.com/NVIDIA/cutlass/tree/main/include/cutlass/gemm/collective/builders) and [convolution](https://github.com/NVIDIA/cutlass/blob/main/include/cutlass/conv/collective/builders/sm100_umma_builder.inl) for the new SM107 [`sm_107a` and `sm_107f` targets](https://github.com/NVIDIA/cutlass/blob/main/CMakeLists.txt): - - Set of unit tests for Rubin SM107 [SIMT GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f32_f32_f32_simt_align1_multi_cluster_shape.cu), [dense FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f8_f8_f8_tensor_op_f32_alignx.cu), [block-scaled FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f8_f8_f8_tensor_op_f32_blockwise.cu), [block-scaled FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f4_f4_f32_tensor_op_f32_2sm_256x192.cu), and [mixed-precision, complex, and 9xBF16 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_umma.cu). + - Set of unit tests for Rubin SM107 [SIMT GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f32_f32_f32_simt_align1_multi_cluster_shape.cu), [dense FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f8_f8_f8_tensor_op_f32_alignx.cu), [block-scaled FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f8_f8_f8_tensor_op_f32_blockwise.cu), [block-scaled FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f4_f4_f32_tensor_op_f32_2sm_256x192.cu), and [mixed-precision, complex, and 9xBF16 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_umma.cu) * Various improvements and fixes from the community and CUTLASS team. Thanks to everyone who submitted PRs! * Optimal code generation with CUDA toolkit versions 13.4. diff --git a/README.md b/README.md index 8e557411fa..e8b8783863 100644 --- a/README.md +++ b/README.md @@ -3,7 +3,7 @@ # CUTLASS 4.8.0 -_CUTLASS 4.8.0 - Aug 2026_ +_CUTLASS 4.8.0 - Sept 2026_ CUTLASS is a collection of abstractions for implementing high-performance matrix-matrix multiplication (GEMM) and related computations at all levels and scales within CUDA. It incorporates strategies for @@ -48,19 +48,21 @@ To get started quickly - please refer : ## CuTe DSL * New features - Initial Rubin support to accelerate dense GEMMs. The following features are available: - - CuTe DSL - - Rubin new FP8 and FP4 Tensor Core support + - CuTe DSL and CuTe extensions + - Support for higher-throughput FP8 (MMA_K=64) and FP4 (MMA_K=128) Tensor Core MMA instructions - B collector reuse - Extended TMEM size from 512 COL to 576 COL - Larger shared memory allocations (328KB) - Enhanced mixed precision throughput (FP8/FP4) + - Softmax acceleration related features - Primitives - - Rubin new FP8 and FP4 Tensor Core support + - Support for higher-throughput FP8 (MMA_K=64) and FP4 (MMA_K=128) Tensor Core MMA instructions + - B collector reuse - Extended TMEM size from 512 COL to 576 COL - - NOTE: Executing Rubin kernels (SM107) requires the R615 driver which will be released - with CUDA Toolkit 13.4 GA. R610 from CUDA Toolkit 13.4 Developer Preview is not - sufficient. + - Larger shared memory allocations (328KB) + - Enhanced mixed precision throughput (FP8/FP4) + - Softmax acceleration related feature + - 2:4 sparsity support for FP4 - CuTe DSL extensions has several new features: - CTA-V maps are now inferred automatically for `cute_ext` TMA load, store, multicast, and reduce-store operations. Explicit CTA-V maps remain supported as overrides. @@ -70,7 +72,7 @@ To get started quickly - please refer : - Improved device-side TMA descriptor updates and grouped GEMM performance through SMEM-staged updates, workspace reuse, and reduced prologue and synchronization overhead. - This release includes an opt-in preview of the CuTe DSL extensions (`cute_ext`) compiler pipeline for ordinary Cute DSL kernels. This pipeline lets user mix `cute_ext` APIs directly into `@cute.jit` and `@cute.kernel` code and is required for kernels that mix the two API surfaces. You may test this feature with the following: `CUTE_DSL_USE_EXTENSION_COMPILER=1 python your_program.py` -The pipeline is expected to preserve program behavior, but generated PTX/SASS may differ. The pipeline is planned to become the default in a future release. +The pipeline is expected to preserve program behavior and performance, but generated PTX/SASS may differ. Note that this pipeline will become the default in the future, no earlier than 4.10. - Added examples for better control over Primitives' compiler warnings/errors introduced in 4.7.0. See the `CuTeDSL/experimental/compiler_diagnostic/` directory. - IKET Profiler Tool - Rubin kernels (sm107) can now be profiled. @@ -84,7 +86,7 @@ The pipeline is expected to preserve program behavior, but generated PTX/SASS ma - Grouped blockscaled GEMM with B collector reuse as applicable - Blockwise GEMM - Rubin (CuTe extension): - - FP4 blockscaled GEMM + - Support for higher-throughput FP8 (MMA_K=64) and FP4 (MMA_K=128) blockscaled GEMM with UE5M3 scale-factor - Grouped GEMM with B collector reuse - Blackwell (CuTe extension): - Dense GEMMs @@ -98,6 +100,11 @@ The pipeline is expected to preserve program behavior, but generated PTX/SASS ma - Input transform GEMM - GeForce pingpong dense GEMM - Blackwell Ultra blockscaled GEMM + - Dense Convolutions + - Implicit-Gemm Fprop Conv + - Blocksclaed Implicit-Gemm Fprop Conv + - GeForce Implicit-Gemm Fprop Conv + - GeForce Blockscaled Implicit-Gemm Fporp Conv - Attention - GQA Decode - Grouped GEMM @@ -105,6 +112,11 @@ The pipeline is expected to preserve program behavior, but generated PTX/SASS ma - Top-K - Ampere (CuTe extension): - SIMT GEMM + - CuTe DSL now supports x86_64 Windows + - CuTe DSL AoT now supports new host target: QNX8.0 + - Notebooks are restructured under examples/python/CuTeDSL/cute/notebooks and new notebooks for primitives will be added under examples/python/CuTeDSL/notebooks + - Numpy is now not a default dependency + * Bug fixes and improvements: - `nvidia-cuda-nvdisasm` is now an optional dependency of `nvidia-cutlass-dsl` via the optional `[sass]` extra. SASS dumping (`CUTE_DSL_KEEP=sass` / KeepSASS) now resolves `nvdisasm` from the bundled wheel (recommended since its version matches the DSL toolchain) or from a local CUDA Toolkit (`CUDA_HOME`/`CUDA_PATH`). A locally-provided `nvdisasm` must come from a CUDA Toolkit at least as new as the toolchain that produced the CUBIN. Installations that never dump SASS are unaffected. - Reduced the protobuf version requirement of IKET profiler from 6.30 to 4.21. This should make protobuf an easier requirement to satisfy in preparation for transioning IKET to an optional extra. @@ -113,49 +125,61 @@ The pipeline is expected to preserve program behavior, but generated PTX/SASS ma vectorized instructions for tensors with a dynamic stride ([!3463](https://github.com/NVIDIA/cutlass/issues/3463)) - Fixed TVM-FFI env stream detection for GPU tensors in tuple ([!3444](https://github.com/NVIDIA/cutlass/issues/3444)) + - Fixed GPU `link-libraries` compile-option order so it is stable across processes + ([!3564](https://github.com/NVIDIA/cutlass/issues/3564)) + - Fixed preprocessor `IndexError` on staged `bool()` with no arguments + ([!3506](https://github.com/NVIDIA/cutlass/issues/3506)) + - Rejected `cute.compile` on `@cute.kernel` with a user error instead of an ICE + ([!3429](https://github.com/NVIDIA/cutlass/issues/3429)) + - Fixed CuTe DSL crashing the Python interpreter when used in a REPL + ([!3413](https://github.com/NVIDIA/cutlass/issues/3413)) + - Fixed a cuDNN Frontend FROST SDPA backward compilation failure issue ([!3594](https://github.com/NVIDIA/cutlass/issues/3594)) + - Fixed SIGABRTs in TVM-FFI launch for cuDNN Frontend SM100 ragged SDPA kernel ([!3595](https://github.com/NVIDIA/cutlass/issues/3595)) This release has been tested against the following packages: - - FlashAttention: [main (0251105)](https://github.com/Dao-AILab/flash-attention/commit/0251105a2fb19d2957484b7f023cd8c115286ced) - - Quack: [main (60d8808)](https://github.com/Dao-AILab/quack/commit/60d88082272a256fa9b3b2ab631c82cfa78337c6) - - FlashInfer: [main (109d44f)](https://github.com/flashinfer-ai/flashinfer/commit/109d44fceea027290d54efcfe927f8a5665b59de) - - cuDNN-Frontend: [deveop (25b3d51)](https://github.com/NVIDIA/cudnn-frontend/commit/25b3d5126b6544afc209e3c2e94a74f5d82db201) - - Pytorch: [main (cf30153)](https://github.com/pytorch/pytorch/commit/cf30153c4c131c8164ee7798e5022d810682e2cb) - - TensorRT-LLM: [main (1cef02e)](https://github.com/NVIDIA/TensorRT-LLM/commit/1cef02e901be43081b1ba6d4981e94ed3bd9c1e8) + - FlashAttention: [main (8d3a3b8)](https://github.com/Dao-AILab/flash-attention/commit/8d3a3b80d4758ebde5a867c50d24d4351443cf2b) + - Quack: [main (35266c3)](https://github.com/Dao-AILab/quack/commit/35266c3298f0e9bf6d5f46c30aace2eaeae517e3) + - FlashInfer: [main (5d0c89e)](https://github.com/flashinfer-ai/flashinfer/commit/5d0c89eacae6ca08f2a1ce92eba557bbad7a1bfc) + - cuDNN-Frontend: [deveop (e0317d1)](https://github.com/NVIDIA/cudnn-frontend/commit/e0317d1f6cbc1bb8e50abce5870b6bf59b9d4a74) + - Pytorch: [main (7d5f021)](https://github.com/pytorch/pytorch/commit/7d5f0216450b8253e62d88284d49de725d88bd42) + - TensorRT-LLM: [main (c295dd9)](https://github.com/NVIDIA/TensorRT-LLM/commit/c295dd9fca143a3fdd2769617699fdd4d4a93eb5) ## CUTLASS Operator API -* Operator API features and functionality: - - Dense and blockscaled GEMMs in Operator API have preliminary Rubin support. These are provided as a preview and may need additional performance tuning. +* Dense and blockscaled GEMMs in Operator API have preliminary Rubin support. These are provided as a preview and may need additional performance tuning. + + - Updated GEMMs include: + + - Dense GEMMs: FP8xFP8 + - Blockscaled GEMM: {MXFP8}x{MXFP4, MXFP8} and {MXFP4, NVFP4}x{MXFP4, NVFP4} (including support for the new UE5M3 scale factor dtype for NVFP4). + + + - These kernels utilize the below new features in Rubin: + + - Higher SMEM (328KB) and TMEM capacity (288KB) + - B-buffer reuse + - Enhanced mixed precision throughput - Updated GEMMs include: - - Dense GEMMs: FP8xFP8 - - Blockscaled GEMM: {MXFP8}x{MXFP4, MXFP8} and {MXFP4, NVFP4}x{MXFP4, NVFP4} (including support for the new UE5M3 scale factor dtype for NVFP4). - - These kernels utilize the below new features in Rubin: - - Higher SMEM (328KB) and TMEM capacity (288KB) - - B-buffer reuse - - Enhanced mixed precision throughput +* Operators can now be ranked by their estimated performance when nvMatmulHeuristics is available. See tutorial here. NOTE: This currently only supports Blackwell kernels as nvMatmulHeuristics does not yet support Rubin. - - Operators can now be ranked by their estimated performance when nvMatmulHeuristics is available. See tutorial [here](https://docs.nvidia.com/cutlass/latest/media/docs/operators/tutorials/007_heuristics.ipynb) +* Standalone kernel implementations are now exposed through `cutlass.kernels`. These kernels can be used directly, in addition to being discoverable and usable via the Operator interface in `cutlass.operators`. - NOTE: This currently only supports Blackwell kernels as nvMatmulHeuristics does not yet support Rubin. +* Custom Epilogue fusions now support per-row or per-column reductions. - - Standalone kernel implementations are now exposed through `cutlass.kernels`, in addition to those exposed through the Operator interface in `cutlass.operators`. This allows kernels to be called directly without looking them up first. - - Custom Epilogue fusions now support partial (per-row or per-column) reductions. - - `IndexPtrGroupedGemmArguments` is now used to represent Grouped GEMM with contiguous-offset/index-pointers. Existing `GroupedGemmArguments` is deprecated and will be removed in a future release. +* `IndexPtrGroupedGemmArguments` is now used to represent Grouped GEMM with contiguous-offset/index-pointers. Existing `GroupedGemmArguments` is deprecated and will be removed in a future release. ## C++ * Added initial Rubin support (SM107) with CuTe C++ building blocks: - - [Rubin Tensor Core MMA instructions](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/mma_sm107_umma.hpp) and corresponding [CuTe MMA traits](https://github.com/NVIDIA/cutlass/blob/main/include/cute/atom/mma_traits_sm107.hpp). + - [Rubin Tensor Core MMA instructions](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/mma_sm107_umma.hpp) and corresponding [CuTe MMA traits](https://github.com/NVIDIA/cutlass/blob/main/include/cute/atom/mma_traits_sm107.hpp). * CuTe examples that demonstrate the use of Rubin SM107 Tensor Core instructions: - - [Dense FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8.cu). - - [Block-scaled FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8_blockscaled.cu). - - [Mixed-precision block-scaled FP8/FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8_blockscaled.cu). - - [Block-scaled FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp4_blockscaled.cu). + - [Dense FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8.cu). + - [Block-scaled FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8_blockscaled.cu). + - [Mixed-precision block-scaled FP8/FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp8_blockscaled.cu). + - [Block-scaled FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/examples/cute/rubin/rubin_fp4_blockscaled.cu). * Adjusted shared-memory and tensor-memory capacity handling for Rubin SM107: - - Set the [SM107 shared-memory capacity](https://github.com/NVIDIA/cutlass/blob/main/include/cutlass/arch/arch.h) to 327 KiB and added launch support for oversized shared-memory configurations. - - Set the [SM107 TMEM capacity](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/tmem_capacity_sm100.hpp) to 576 columns per SM, updated the [CuTe 1SM and 2SM TMEM allocators](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/tmem_allocator_sm100.hpp) for Rubin's exclusive allocation path. + - Set the [SM107 shared-memory capacity](https://github.com/NVIDIA/cutlass/blob/main/include/cutlass/arch/arch.h) to 327 KiB and added launch support for oversized shared-memory configurations. + - Set the [SM107 TMEM capacity](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/tmem_capacity_sm100.hpp) to 576 columns per SM, updated the [CuTe 1SM and 2SM TMEM allocators](https://github.com/NVIDIA/cutlass/blob/main/include/cute/arch/tmem_allocator_sm100.hpp) for Rubin's exclusive allocation path. * Enabled the existing SM100-compatible [GEMM](https://github.com/NVIDIA/cutlass/tree/main/include/cutlass/gemm/collective/builders) and [convolution](https://github.com/NVIDIA/cutlass/blob/main/include/cutlass/conv/collective/builders/sm100_umma_builder.inl) for the new SM107 [`sm_107a` and `sm_107f` targets](https://github.com/NVIDIA/cutlass/blob/main/CMakeLists.txt): - - Set of unit tests for Rubin SM107 [SIMT GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f32_f32_f32_simt_align1_multi_cluster_shape.cu), [dense FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f8_f8_f8_tensor_op_f32_alignx.cu), [block-scaled FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f8_f8_f8_tensor_op_f32_blockwise.cu), [block-scaled FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f4_f4_f32_tensor_op_f32_2sm_256x192.cu), and [mixed-precision, complex, and 9xBF16 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_umma.cu). + - Set of unit tests for Rubin SM107 [SIMT GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f32_f32_f32_simt_align1_multi_cluster_shape.cu), [dense FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f8_f8_f8_tensor_op_f32_alignx.cu), [block-scaled FP8 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f8_f8_f8_tensor_op_f32_blockwise.cu), [block-scaled FP4 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_f4_f4_f32_tensor_op_f32_2sm_256x192.cu), and [mixed-precision, complex, and 9xBF16 GEMM](https://github.com/NVIDIA/cutlass/blob/main/test/unit/gemm/device/sm107_gemm_umma.cu) Note: Executing Rubin kernels (SM107) requires the R615 driver which will be released with CUDA Toolkit 13.4 GA. R610 from CUDA Toolkit 13.4 Developer Preview is not sufficient. diff --git a/examples/python/CuTeDSL/cute/blackwell/kernel/blockwise_gemm/blockwise_gemm.py b/examples/python/CuTeDSL/cute/blackwell/kernel/blockwise_gemm/blockwise_gemm.py index 3059e8f342..b997bd6ee1 100644 --- a/examples/python/CuTeDSL/cute/blackwell/kernel/blockwise_gemm/blockwise_gemm.py +++ b/examples/python/CuTeDSL/cute/blackwell/kernel/blockwise_gemm/blockwise_gemm.py @@ -33,7 +33,7 @@ import cutlass import cutlass.cute as cute -import cutlass.cute.testing as testing +from cutlass import testing from cutlass.cute.nvgpu import cpasync, tcgen05 import cutlass.utils as utils import cutlass.pipeline as pipeline @@ -78,7 +78,7 @@ .. code-block:: bash - python examples/blackwell/blockwise_gemm/blockwise_gemm.py \ + python examples/cute/blackwell/kernel/blockwise_gemm/blockwise_gemm.py \ --ab_dtype Float8E4M3FN --c_dtype BFloat16 --acc_dtype Float32 \ --scale_dtype Float32 \ --mma_tiler_mn 128,128 --cluster_shape_mn 1,2 \ @@ -88,7 +88,7 @@ .. code-block:: bash - ncu python examples/blackwell/blockwise_gemm/blockwise_gemm.py \ + ncu python examples/cute/blackwell/kernel/blockwise_gemm/blockwise_gemm.py \ --ab_dtype Float8E4M3FN --c_dtype BFloat16 --acc_dtype Float32 \ --scale_dtype Float32 \ --mma_tiler_mn 128,128 --cluster_shape_mn 1,2 \ @@ -233,7 +233,7 @@ def __init__( barrier_id=3, num_threads=self.threads_per_warp, ) - self.num_smem_capacity = utils.get_smem_capacity_in_bytes("sm_100") + self.num_smem_capacity = cutlass.memory.get_smem_capacity_in_bytes("sm_100") # TMEM offset for final accumulator self.tmem_final_offset = 384 @@ -254,6 +254,7 @@ def _setup_attributes(self): # Configure tiled mma tiled_mma = sm100_utils.make_trivial_tiled_mma( self.a_dtype, + self.b_dtype, self.a_major_mode, self.b_major_mode, self.acc_dtype, @@ -422,9 +423,13 @@ def __call__( self.c_dtype: Type[cutlass.Numeric] = c.element_type self.sfa_dtype: Type[cutlass.Numeric] = sfa.element_type self.sfb_dtype: Type[cutlass.Numeric] = sfb.element_type - self.a_major_mode = utils.LayoutEnum.from_tensor(a).mma_major_mode() - self.b_major_mode = utils.LayoutEnum.from_tensor(b).mma_major_mode() - self.c_layout = utils.LayoutEnum.from_tensor(c) + self.a_major_mode = cutlass.tensor_utils.LayoutEnum.from_tensor( + a + ).mma_major_mode() + self.b_major_mode = cutlass.tensor_utils.LayoutEnum.from_tensor( + b + ).mma_major_mode() + self.c_layout = cutlass.tensor_utils.LayoutEnum.from_tensor(c) # Check if input data types are compatible with MMA instruction if cutlass.const_expr(self.a_dtype != self.b_dtype): @@ -434,6 +439,7 @@ def __call__( self._setup_attributes() tiled_mma = sm100_utils.make_trivial_tiled_mma( + self.a_dtype, self.a_dtype, self.a_major_mode, self.b_major_mode, @@ -678,15 +684,12 @@ def kernel( # # Alloc and init: a+b full/empty, accumulator full/empty, tensor memory dealloc barrier # - smem = utils.SmemAllocator() + smem = cutlass.memory.SmemAllocator() storage = smem.allocate(self.shared_storage) # Initialize mainloop ab_pipeline (barrier) and states ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) - num_tma_producer = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1 - ab_pipeline_consumer_group = pipeline.CooperativeGroup( - pipeline.Agent.Thread, num_tma_producer - ) + ab_pipeline_consumer_group = pipeline.CooperativeGroup(pipeline.Agent.Warp) ab_pipeline = pipeline.PipelineTmaUmma.create( barrier_storage=storage.ab_mbar_ptr.data_ptr(), num_stages=self.num_ab_stage, @@ -694,6 +697,7 @@ def kernel( consumer_group=ab_pipeline_consumer_group, tx_count=self.num_tma_load_bytes, cta_layout_vmnk=cluster_layout_vmnk, + enable_multicast_signaling=True, defer_sync=True, ) @@ -766,7 +770,7 @@ def kernel( ) # Tensor memory dealloc barrier init - tmem = utils.TmemAllocator( + tmem = cutlass.memory.TmemAllocator( storage.tmem_holding_buf.ptr, barrier_for_retrieve=self.tmem_alloc_barrier, allocator_warp_id=self.epilog_warp_id[0], @@ -1401,24 +1405,10 @@ def kernel( ) # tCtAcc += tCrA * tCrB - num_kblocks = cute.size(tCrA, mode=[2]) - for kblock_idx in cutlass.range(num_kblocks, unroll_full=True): - kblock_coord = ( - None, - None, - kblock_idx, - ab_consumer_state.index, - ) - - cute.gemm( - tiled_mma, - tCtAcc, - tCrA[kblock_coord], - tCrB[kblock_coord], - tCtAcc, - ) - # Enable accumulate on tCtAcc after first kblock - tiled_mma.set(tcgen05.Field.ACCUMULATE, True) + tile_crd = (None, None, None, ab_consumer_state.index) + cute.gemm( + tiled_mma, tCtAcc, tCrA[tile_crd], tCrB[tile_crd], tCtAcc + ) # Async arrive AB buffer empty ab_pipeline.consumer_release(ab_consumer_state) @@ -1736,7 +1726,6 @@ def kernel( tTR_rC = None tiled_copy_r2s = None - simt_atom = None tRS_rC = None tRS_sC = None bSG_sC = None @@ -2206,7 +2195,7 @@ def _compute_stages( b_dtype: Type[cutlass.Numeric], epi_tile: cute.Tile, c_dtype: Type[cutlass.Numeric], - c_layout: utils.LayoutEnum, + c_layout: cutlass.tensor_utils.LayoutEnum, sfa_dtype: Type[cutlass.Numeric], sfb_dtype: Type[cutlass.Numeric], sfa_count: int, @@ -2229,7 +2218,7 @@ def _compute_stages( :param c_dtype: Data type of operand C (output). :type c_dtype: type[cutlass.Numeric] :param c_layout: Layout of operand C. - :type c_layout: utils.LayoutEnum + :type c_layout: cutlass.tensor_utils.LayoutEnum :param num_smem_capacity: Total available shared memory capacity in bytes. :type num_smem_capacity: int :param occupancy: Target number of CTAs per SM (occupancy). @@ -2672,7 +2661,7 @@ def run( b_major, c_major, ): - raise TypeError( + raise cutlass.testing.CantImplementError( f"Unsupported testcase {ab_dtype}, {acc_dtype}, {c_dtype}, {use_2cta_instrs}, {mma_tiler_mn}, {cluster_shape_mn}, {m}, {n}, {k}, {l}, {a_major}, {b_major}, {c_major}" ) @@ -2718,10 +2707,7 @@ def run( current_stream, options=( "--opt-level=2" - if ( - cutlass.target_version(exact_version="12.9") - or cutlass.target_version(min_version="13.1") - ) + if tuple(map(int, cutlass.__version__.split(".")[:2])) >= (4, 6) else "" ), ) diff --git a/examples/python/CuTeDSL/cute/blackwell/kernel/conv/dense_blockscaled_implicit_gemm_fprop.py b/examples/python/CuTeDSL/cute/blackwell/kernel/conv/dense_blockscaled_implicit_gemm_fprop.py new file mode 100644 index 0000000000..7fbfc60a2d --- /dev/null +++ b/examples/python/CuTeDSL/cute/blackwell/kernel/conv/dense_blockscaled_implicit_gemm_fprop.py @@ -0,0 +1,5369 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause + +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: + +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. + +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. + +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. + +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import argparse +from typing import Optional, Tuple, Type, Union, Literal +from functools import lru_cache +import cuda.bindings.driver as cuda +import sys +import os + +import torch +import torch.nn.functional as F + +import cutlass +import cutlass.cute as cute +from cutlass import testing +from cutlass.cute.runtime import from_dlpack +import cutlass.torch as cutlass_torch +import cutlass.utils as utils +from cutlass.cute.nvgpu import cpasync, tcgen05 +import cutlass.pipeline as pipeline +from dataclasses import dataclass +from cutlass.pipeline import ( + Agent, + CooperativeGroup, + PipelineOp, + PipelineState, + PipelineAsync, + agent_sync, + pipeline_init_arrive, + pipeline_init_wait, +) +import cutlass.utils.blackwell_helpers as sm100_utils +import cutlass.utils.blockscaled_layout as blockscaled_utils +from pathlib import Path + +if __name__ == "__main__": + current_dir = os.path.dirname(os.path.abspath(__file__)) + sys.path.insert(0, os.path.join(current_dir, "../../../")) + # `helpers` sits at the examples/CuTeDSL root; running this file as a + # script only puts its own directory on sys.path. + cutedsl_dir = str(Path(__file__).resolve().parents[4]) + if cutedsl_dir not in sys.path: + sys.path.insert(0, cutedsl_dir) + +from blackwell.kernel.conv.dense_implicit_gemm_fprop import ( + Sm100PersistentDenseImplicitGemmFpropKernel, + _check_im2col_descriptor_limits, + _check_swizzle_size, + _check_tensor_alignment, + _parse_comma_separated_ints, + _rt_conv_scalars, + compute_zpq, + prepare_tensors, +) + +from cutlass.cutlass_dsl import if_generate as _if_generate +from cutlass._mlir.dialects import nvvm +from cutlass.memory import SmemAllocator, TmemAllocator +from cutlass.tensor_utils import LayoutEnum +from cutlass.utils import ( + ClcDynamicPersistentTileScheduler, + ClcDynamicPersistentTileSchedulerParams, +) +from cutlass.memory import get_num_tmem_alloc_cols, get_smem_capacity_in_bytes + + +@dataclass(frozen=True) +class PipelineTmaCpAsyncUmma(PipelineAsync): + """ + Merged pipeline for TMA (A+B+SFB) producers AND cp.async (SFA) producers sharing a single mbarrier set. + + Full mbarrier uses BOTH independent mbarrier counters: + - tx_count : incremented via mbarrier.expect_tx for TMA bytes, decremented by TMA completion + - arrive_cnt : initialized to cp.async producer group size; decremented by cp.async.mbarrier.arrive.noinc + + Barrier phase flips only when tx_count == 0 AND arrive_cnt == 0, so the UMMA consumer performs + a single wait/release per stage covering both producer streams. + + General variant: supports arbitrary cluster shapes (including cluster_shape_n > 1). + """ + + is_leader_cta: bool + cta_group: cute.nvgpu.tcgen05.CtaGroup + tx_count: int + + @staticmethod + def create( + *, + num_stages: int, + cpasynd_producer_group: CooperativeGroup, + consumer_group: CooperativeGroup, + tx_count: int, + barrier_storage: cute.Pointer, + cta_layout_vmnk: Optional[cute.Layout] = None, + mcast_mode_mn: Tuple[int, int] = (1, 1), + defer_sync: bool = False, + ) -> "PipelineTmaCpAsyncUmma": + """Creates and initializes a merged TMA+cp.async / UMMA pipeline.""" + if not isinstance(barrier_storage, cute.Pointer): + raise ValueError( + f"Expected barrier_storage to be a cute.Pointer, but got {type(barrier_storage)}" + ) + + producer = (PipelineOp.AsyncLoad, cpasynd_producer_group) + consumer = (PipelineOp.TCGen05Mma, consumer_group) + + sync_object_full = pipeline.PipelineTmaUmma._make_sync_object( + barrier_storage.align(min_align=8), num_stages, producer, tx_count + ) + sync_object_empty = pipeline.PipelineTmaUmma._make_sync_object( + barrier_storage.align(min_align=8) + num_stages, num_stages, consumer + ) + + cta_group = ( + cute.nvgpu.tcgen05.CtaGroup.ONE + if cta_layout_vmnk is None or cute.size(cta_layout_vmnk, mode=[0]) == 1 + else cute.nvgpu.tcgen05.CtaGroup.TWO + ) + + if cta_layout_vmnk is None or cute.size(cta_layout_vmnk) == 1: + consumer_mask = None + is_leader_cta = True + # No cross-CTA; remote arrives degenerate to self-arrive. + producer_mask = cutlass.Int32(0) + else: + # Compute the empty-drain multicast mask: OR over A-axis and + # B-axis TMA multicast sets (plus their V-peers). + # PipelineTmaUmma._compute_mcast_arrival_mask returns the same + # value but invoking it here produces stale IR for the merged + # sync_object_empty; inline the equivalent computation instead. + cta_rank_here = cute.arch.make_warp_uniform( + cute.arch.block_idx_in_cluster() + ) + coord_self = cta_layout_vmnk.get_flat_coord(cta_rank_here) + coord_peer = (coord_self[0] ^ 1, *coord_self[1:]) + # SFA cp.async remote-arrive target. Only meaningful for 2CTA, where + # the peer CTA's cp.async lands in the leader's SMEM and both CTAs + # must remote-arrive on the leader's sfa_full to close the peer-sSFA + # race. The 2CTA flat rank is v + m*V + ..., so clearing the V bit + # yields the leader rank. In a non-2CTA cluster (V=1) every CTA owns + # its SFA locally, so producer_arrive_remote stays on the local + # barrier and this value is unused; keep it as self-rank. + producer_mask = ( + cta_rank_here & ~cutlass.Int32(1) + if cta_group == cute.nvgpu.tcgen05.CtaGroup.TWO + else cta_rank_here + ) + mask_a_self = cute.nvgpu.cpasync.create_tma_multicast_mask( + cta_layout_vmnk, coord_self, mcast_mode=2 + ) + mask_b_self = cute.nvgpu.cpasync.create_tma_multicast_mask( + cta_layout_vmnk, coord_self, mcast_mode=1 + ) + mask_a_peer = cute.nvgpu.cpasync.create_tma_multicast_mask( + cta_layout_vmnk, coord_peer, mcast_mode=2 + ) + mask_b_peer = cute.nvgpu.cpasync.create_tma_multicast_mask( + cta_layout_vmnk, coord_peer, mcast_mode=1 + ) + if mcast_mode_mn[0] == 1 and mcast_mode_mn[1] == 1: + consumer_mask = ( + cutlass.Int32(mask_a_self) + | cutlass.Int32(mask_b_self) + | cutlass.Int32(mask_a_peer) + | cutlass.Int32(mask_b_peer) + ) + elif mcast_mode_mn[1] == 1: + consumer_mask = cutlass.Int32(mask_b_self) | cutlass.Int32(mask_b_peer) + else: + assert mcast_mode_mn[0] == 1 + consumer_mask = cutlass.Int32(mask_a_self) | cutlass.Int32(mask_a_peer) + is_leader_cta = pipeline.PipelineTmaUmma._compute_is_leader_cta( + cta_layout_vmnk + ) + + if not defer_sync: + cute.arch.mbarrier_init_fence() + if cta_layout_vmnk is None or cute.size(cta_layout_vmnk) == 1: + agent_sync(Agent.ThreadBlock) + else: + agent_sync(Agent.ThreadBlockCluster, is_relaxed=True) + + return PipelineTmaCpAsyncUmma( + sync_object_full, + sync_object_empty, + num_stages, + producer_mask, + consumer_mask, + is_leader_cta, + cta_group, + tx_count, + ) + + def producer_acquire_tma( + self, + state: PipelineState, + try_acquire_token: Optional[cutlass.Boolean] = None, + *, + expected_tx: Optional[cutlass.Int32] = None, + ) -> None: + """TMA producer path. + + Only the leader CTA issues ``arrive_and_expect_tx``. The multicast + variant broadcasts BOTH the arrive and the tx-expect to every peer + barrier in the cluster, so peer CTAs must do nothing here. Any peer + ``arrive(1)`` would double-count the arrive on the peer barrier and + opens a race window against cp.async producers draining first. + """ + _if_generate( + try_acquire_token is None or try_acquire_token == 0, + lambda: self.sync_object_empty.wait(state.index, state.phase), + ) + tx = self.tx_count if expected_tx is None else expected_tx + + def _leader_arrive_and_expect_tx() -> None: + self.sync_object_full.arrive_and_expect_tx(state.index, tx) + + _if_generate( + self.is_leader_cta, + _leader_arrive_and_expect_tx, + ) + + def consumer_release(self, state: PipelineState) -> None: + """UMMA consumer releases the shared empty barrier once per stage (multicast to peer CTAs).""" + self.sync_object_empty.arrive(state.index, self.consumer_mask, self.cta_group) + + # ------------------------------------------------------------------ + # Shift-by-K cross-CTA producer commit primitives. + # In 2CTA clusters, cp.async.mbarrier.arrive.noinc only signals the + # local CTA's sfa_full, so the peer's cp.async landing on leader SMEM + # is not covered. Split the commit into commit_group + wait_group(K) + # + cross-CTA mbarrier.arrive(dst=leader_rank) so that arrives are + # deferred by K iters, giving cp.async time to land on both CTAs + # before the leader MMA fires. + def producer_cp_async_commit(self) -> None: + """Register all prior cp.async copies as one cp.async group.""" + cute.arch.cp_async_commit_group() + + def producer_cp_async_wait(self, num_inflight: int) -> None: + """Block until at most ``num_inflight`` cp.async groups remain pending.""" + cute.arch.cp_async_wait_group(num_inflight) + + def producer_arrive_remote(self, stage_index) -> None: + """Signal SFA cp.async completion on sfa_full[stage_index]. + + In a 2CTA cluster the peer CTA's cp.async lands in the leader's SMEM, + so both CTAs must arrive on the leader's sfa_full via DSMEM (producer_mask + holds the leader rank). In a non-2CTA cluster every CTA owns its SFA + locally, so the arrive stays on the local barrier with no DSMEM + retargeting. + """ + bar = self.sync_object_full.get_barrier(stage_index) + if cutlass.const_expr(self.cta_group == cute.nvgpu.tcgen05.CtaGroup.TWO): + cute.arch.mbarrier_arrive(bar, self.producer_mask) + else: + cute.arch.mbarrier_arrive(bar) + + +""" +A high-performance 3D implicit-GEMM based fprop convolution example for the NVIDIA Blackwell SM100 architecture using CUTE DSL. +- Input tensor A is NxDxHxWxC, must be C major. +- Filter tensor B is KxTxRxSxC, must be C major. +- Output tensor D is NxZxPxQxK, must be K major. + +This kernel supports the following features: + - Utilizes Tensor Memory Access (TMA) for efficient memory operations and for the im2col transformation of the input tensor A. + - Utilizes Blackwell's tcgen05 MMA for matrix multiply-accumulate (MMA) operations (including 2cta mma instructions) + - Implements TMA multicast with cluster to reduce L2 memory traffic + - Utilizes either a dense GEMM kernel or a persistent dense GEMM kernel for the implicit GEMM. + +This implicit-GEMM based convolution works by converting the convolution into a GEMM problem with the following mapping: +- GEMM M dimension maps to NxZxPxQ +- GEMM N dimension maps to K +- GEMM K dimension maps to TxRxSxC +During the load of input tensor to SMEM, the TMA operation performs the im2col transformation on the input tensor A. +This transforms the A matrix into the required shape for the GEMM operation (NxZxPxQ by TxRxSxC), and may involve replication of the input elements. +Filter tensor can be loaded to SMEM without any transformation. +The output tensor D is then stored to GMEM via TMA im2col store (no transformation necessary). + +To run this example: + +.. code-block:: bash + + python examples/CuTeDSL/cute/blackwell/kernel/conv/dense_blockscaled_implicit_gemm_fprop.py \ + --ncdhw 1,128,32,32,32 --ktrs 256,3,3,3 \ + --use_2cta_instrs --mma_tiler_mn 256,128 \ + --preferred_cluster_shape_mn 2,1 --fallback_cluster_shape_mn 2,1 \ + --upper_pad_dhw 1,1,1 --lower_pad_dhw 1,1,1 \ + --stride_dhw 1,1,1 --dil_dhw 1,1,1 + +To collect performance with NCU profiler: + +.. code-block:: bash + + ncu python examples/CuTeDSL/cute/blackwell/kernel/conv/dense_blockscaled_implicit_gemm_fprop.py \ + --ncdhw 1,128,32,32,32 --ktrs 256,3,3,3 \ + --use_2cta_instrs --mma_tiler_mn 256,128 \ + --preferred_cluster_shape_mn 2,1 --fallback_cluster_shape_mn 2,1 \ + --upper_pad_dhw 1,1,1 --lower_pad_dhw 1,1,1 \ + --stride_dhw 1,1,1 --dil_dhw 1,1,1 \ + --warmup_iterations 1 --iterations 10 --skip_ref_check + +Constraints: +* A/B data type is Float4E2M1FN (NVFP4) or Float8E4M3FN (MXFP8). Each pins its + scale-factor format: NVFP4 takes a Float8E4M3FN scale over 16 channels, MXFP8 a + Float8E8M0FNU scale over 32. A block-scaled narrow output takes the same format + for its own scale factor, which pairs NVFP4 with an FP4 output and MXFP8 with an + FP8 E4M3 one; a 16-bit output carries no scale factor. +* A/B tensor must have the same data type +* Mma tiler M must be 128 with use_2cta_instrs=False or 256 with True, the two + shapes giving the per-CTA M of 128 that the block-scaled MMA requires +* Mma tiler N must be a multiple of 32 up to 256. Only 64, 128, 192 and 256 have a + per-tile offset into the staged scale-factor chunk, so the others serve problems + needing a single N tile +* Cluster shape M/N must be positive and power of 2, total cluster size <= 16 +* Cluster shape M must be multiple of 2 if use_2cta_instrs=True +* The contiguous dimension of A/B/D tensors must be at least 16 bytes aligned, + i.e, number of elements is a multiple of 4, 8, and 16 for TFloat32, + Float16/BFloat16, and Int8/Uint8/Float8, respectively. +* The GEMM-K tile is an independent compile-time choice: 64, 128 or 256 channels + for NVFP4 (one, two or four MMA K instructions), 128 for MXFP8. It need not + divide the input channel count C, which only has to be a whole number of + scale-factor atoms (4 * sf_vec_size channels). A K tile wider than the channels + a filter position has left is a partial tile: the im2col TMA zero-fills the A + and B channels past C, the SFA transfer is clamped back inside the pixel's + scale-factor row, and SFB is addressed per filter position over a span rounded + out to whole K tiles, so the tail contributes nothing. +""" + + +class Sm100BlockScaledPersistentDenseImplicitGemmFpropKernel( + Sm100PersistentDenseImplicitGemmFpropKernel +): + """ + Persistent 3D convolution kernel. + The input (A) is expected to be in 5D tensor (NDHWC) format and is loaded via TMA im2col load atom. + The filter (B) is expected to be in 5D tensor (KTRSC) format and is loaded via TMA load atom. + The output (D) is expected to be in 5D tensor (NZPQK) format and is stored via TMA im2col store. + The implicit GEMM runs as a megakernel that dispatches on the cluster shape the + launch actually formed, so a single launch serves both the preferred and the + fallback cluster. + + :param acc_dtype: Data type for accumulation during computation + :type acc_dtype: type[cutlass.Numeric] + :param use_2cta_instrs: Whether to use CTA group 2 for advanced thread cooperation + :type use_2cta_instrs: bool + :param mma_tiler_mn: Shape of the Matrix Multiply-Accumulate (MMA) tiler (M,N) + :type mma_tiler_mn: Tuple[int, int] + :param cluster_shape_mn: Cluster dimensions (M,N) for parallel processing + :type cluster_shape_mn: Tuple[int, int] + :param filter_trs: Filter dimensions (T, R, S) + :type filter_trs: Tuple[int, int, int] + :param upper_padding_dhw: Upper padding (PadD, PadH, PadW) + :type upper_padding_dhw: Tuple[int, int, int] + :param lower_padding_dhw: Lower padding (PadD, PadH, PadW) + :type lower_padding_dhw: Tuple[int, int, int] + :param stride_dhw: Stride (Sd, Sh, Sw) + :type stride_dhw: Tuple[int, int, int] + :param dilation_dhw: Dilation (DilD, DilH, DilW) + :type dilation_dhw: Tuple[int, int, int] + :param swizzle_size: Swizzling size in the unit of cluster. 1 means no swizzle + :type swizzle_size: int + :param raster_along: Rasterization order of clusters. Only used when swizzle_size > 1. + :type raster_along: Literal["m", "n"] + + :note: In current version, A and B tensor must be C major. D tensor must be K major. + + :note: Supported A/B data types, which A and B share: + - Float4E2M1FN (NVFP4), scaled by a Float8E4M3FN factor over 16 channels + - Float8E4M3FN (MXFP8), scaled by a Float8E8M0FNU factor over 32 channels + + :note: The accumulator is Float32. + + :note: Supported D data types: + - the block-scaled narrow output that pairs with A/B -- Float4E2M1FN for + NVFP4, Float8E4M3FN for MXFP8 -- which carries its own scale factor + - Float16/BFloat16, which carry none + + :note: Constraints: + - MMA tiler M must be 128 (use_2cta_instrs=False) or 256 (use_2cta_instrs=True), + the two shapes giving the per-CTA M of 128 the block-scaled MMA requires + - MMA tiler N must be 64, 128, 192 or 256; the other multiples of 32 up to 256 + serve problems needing a single N tile + - Cluster shape M must be multiple of 2 if use_2cta_instrs=True + - Cluster shape M/N must be positive and power of 2, total cluster size <= 16 + - C must be a whole number of scale-factor atoms, 4 * sf_vec_size channels + + **Example:** + + .. code-block:: python + conv = Sm100BlockScaledPersistentDenseImplicitGemmFpropKernel( + acc_dtype=cutlass.Float32, + sf_vec_size=16, + use_2cta_instrs=True, + mma_tiler_mn=(256, 128), + preferred_cluster_shape_mn=(4, 2), + fallback_cluster_shape_mn=(2, 1), + cta_tile_k=128, + ) + conv( + a, b, d, sfa, sfb, alpha, epilogue_op, + bias_tensor=bias, # optional length-K bias tensor, or None + rt_upper_pad_d=upper_pad_d, + rt_upper_pad_h=upper_pad_h, + rt_upper_pad_w=upper_pad_w, + rt_lower_pad_d=lower_pad_d, + rt_lower_pad_h=lower_pad_h, + rt_lower_pad_w=lower_pad_w, + rt_stride_d=stride_d, + rt_stride_h=stride_h, + rt_stride_w=stride_w, + rt_dil_d=dil_d, + rt_dil_h=dil_h, + rt_dil_w=dil_w, + stream=stream, + ) + """ + + def __init__( + self, + acc_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + use_2cta_instrs: bool, + mma_tiler_mn: Tuple[int, int], + preferred_cluster_shape_mn: Tuple[int, int], + fallback_cluster_shape_mn: Tuple[int, int], + cta_tile_k: int, + swizzle_size: int = 1, + raster_along: Literal["m", "n"] = "m", + ): + super().__init__( + acc_dtype=acc_dtype, + use_2cta_instrs=use_2cta_instrs, + mma_tiler_mn=mma_tiler_mn, + preferred_cluster_shape_mn=preferred_cluster_shape_mn, + fallback_cluster_shape_mn=fallback_cluster_shape_mn, + swizzle_size=swizzle_size, + raster_along=raster_along, + ) + + self.sf_vec_size = sf_vec_size + + # cta_tile_k is the compile-time K tile (shapes the SMEM allocation). + # The input channel count C is read from the dynamic tensors at runtime and + # bounded only by the alignment of the channel axis, so one cubin serves + # every C this kernel accepts. All output geometry (N/Z/P/Q/K) is runtime + # as well. + self.cta_tile_k = cta_tile_k + + # Override warp specialization for fp4 conv: 4 cp.async SFA warps + sched warp. + self.epilogue_warp_id = (0, 1, 2, 3) + self.mma_warp_id = 4 + self.tma_warp_id = 5 + self.cpasync_sfa_warp_id = (6, 7, 8, 9) + self.sched_warp_id = 10 + # Dedicated residual G2S warp, only spawned when beta != 0 (has_residual). + # threads_per_cta is finalized in __call__ once has_residual is known. + self.residual_warp_id = 11 + self.base_warp_ids = ( + self.mma_warp_id, + self.tma_warp_id, + *self.cpasync_sfa_warp_id, + self.sched_warp_id, + *self.epilogue_warp_id, + ) + self.threads_per_cta = 32 * len(self.base_warp_ids) + # Override barrier ids; parent's preferred_cluster init reserves bar 0 for cta_sync. + self.epilog_sync_bar_id = 1 + self.epilog_sync_barrier = pipeline.NamedBarrier( + barrier_id=self.epilog_sync_bar_id, + num_threads=32 * len(self.epilogue_warp_id), + ) + self.tmem_alloc_sync_bar_id = 2 + + def _setup_conv_input_attrs(self, a_tensor, b_tensor, d_tensor): + """Validate and set input-dependent attributes. + + Sets a_dtype, b_dtype, d_dtype, a_major_mode, b_major_mode, d_layout. + """ + self.a_dtype: Type[cutlass.Numeric] = a_tensor.element_type + self.b_dtype: Type[cutlass.Numeric] = b_tensor.element_type + self.d_dtype: Type[cutlass.Numeric] = d_tensor.element_type + if cutlass.const_expr(a_tensor.leading_dim != 4): + raise RuntimeError("The layout of a_tensor is not supported") + if cutlass.const_expr(b_tensor.leading_dim != 4): + raise RuntimeError("The layout of b_tensor is not supported") + if cutlass.const_expr(d_tensor.leading_dim != 4): + raise RuntimeError("The layout of d_tensor is not supported") + self.a_major_mode = cute.nvgpu.OperandMajorMode.K + self.b_major_mode = cute.nvgpu.OperandMajorMode.K + self.d_layout = LayoutEnum.ROW_MAJOR + + if cutlass.const_expr(self.a_dtype != self.b_dtype): + raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}") + + # Block-scaled D (SFD) output. Caller opts in by passing sfd_tensor; + # gen_sfd is decided in __call__ from it. A narrow output carries one scale + # factor per sfd_vec_size elements; a 16-bit output carries none. The SFD + # block is the input SF block -- same element type (bound in + # __call__) and same vector size -- so a quantized output feeds the next + # operation directly: NVFP4 is E4M3 over 16 channels, MXFP8 E8M0 over 32. + self.sfd_vec_size: int = self.sf_vec_size + # M_D is the largest absolute value representable in the output dtype, + # the full-scale target when quantizing each block to it. FP4 (E2M1) + # magnitudes are {0,.5,1,1.5,2,3,4,6} so max is 6.0; E4M3 max normal is + # 448.0. Only narrow SFD outputs use it, hence None otherwise. + self.M_D: Optional[float] = { + cutlass.Float4E2M1FN: 6.0, + cutlass.Float8E4M3FN: 448.0, + }.get(self.d_dtype, None) + + def _setup_conv_tma( + self, + a_tensor, + b_tensor, + d_tensor, + c_tensor, + tiled_mma, + upper_pad_op, + lower_pad_op, + stride_op, + dil_op, + ): + """Set up dual-cluster TMA atoms and tensors for im2col convolution. + + Creates preferred and fallback TMA atoms for A and B tensors, and a single + TMA atom for D (im2col store is cluster-independent). + + upper_pad_op/lower_pad_op/stride_op/dil_op are the pad/stride/dilation + operand tuples that build the im2col A descriptor corner. They are + runtime Int32 tuples so one compiled cubin serves any pad/stride/dil + config without recompilation. + + :returns: (tma_atom_a_preferred, tma_tensor_a_preferred, + tma_atom_a_fallback, tma_tensor_a_fallback, + tma_atom_b_preferred, tma_tensor_b_preferred, + tma_atom_b_fallback, tma_tensor_b_fallback, + tma_atom_d, tma_tensor_d) + """ + + ( + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + ) = self._make_a_im2col_tma_atoms( + a_tensor, b_tensor, tiled_mma, upper_pad_op, lower_pad_op, stride_op, dil_op + ) + + # --- B: tiled TMA load --- + b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) + # Filter (K, T, R, S, C) -> reorder to (K, (C, S, R, T)) for coord_iter indexing. + mB = cute.make_tensor( + b_tensor.iterator, cute.select(b_tensor.layout, mode=[0, 4, 3, 2, 1]) + ) + mB = cute.group_modes(mB, begin=1, end=5) + b_internal_type = ( + cutlass.TFloat32 if mB.element_type is cutlass.Float32 else None + ) + + # --- B preferred --- + b_op_preferred = sm100_utils.cluster_shape_to_tma_atom_B( + self.preferred_cluster_shape_mn, tiled_mma.thr_id + ) + tma_atom_b_preferred, tma_tensor_b_preferred = cute.nvgpu.make_tiled_tma_atom_B( + b_op_preferred, + mB, + b_smem_layout, + self.mma_tiler, + tiled_mma, + self.preferred_cluster_layout_vmnk.shape, + internal_type=b_internal_type, + ) + + # --- B fallback --- + b_op_fallback = sm100_utils.cluster_shape_to_tma_atom_B( + self.fallback_cluster_shape_mn, tiled_mma.thr_id + ) + tma_atom_b_fallback, tma_tensor_b_fallback = cute.nvgpu.make_tiled_tma_atom_B( + b_op_fallback, + mB, + b_smem_layout, + self.mma_tiler, + tiled_mma, + self.fallback_cluster_layout_vmnk.shape, + internal_type=b_internal_type, + ) + + # CLC response size is 4B * 4 elements + self.num_clc_response_bytes = 16 + + # --- D: TMA im2col store (cluster-independent) --- + mD = cute.make_tensor( + d_tensor.iterator, cute.select(d_tensor.layout, mode=[3, 2, 1, 0, 4]) + ) + mD = cute.group_modes(mD, begin=0, end=4) + + epi_smem_layout = cute.slice_( + self.d_smem_layout_staged, (None, None, (None, 0)) + ) + + tma_atom_d, tma_tensor_d = cpasync.make_im2col_tma_atom( + cpasync.CopyBulkTensorIm2ColS2GOp(), + mD, + epi_smem_layout, + self.epi_tile, + ) + tma_tensor_d = cute.coalesce(tma_tensor_d, target_profile=(1, 1)) + + # --- Residual: im2col G2S load, the inverse of the D im2col store --- + # The residual has the same (N,Z,P,Q,K) shape/layout as D, so it reuses + # mD's hierarchical view and D's epilogue smem layout. Load mode (G2S) + # requires the full im2col descriptor corner set; the residual is an + # identity element map (a 1x1x1 tap with no padding/stride/dilation), so + # every corner/stride is trivial: lower/upper corners and padding are 0, + # the DHW stride is 1, and the SRT lower/stride are 0/1. + if cutlass.const_expr(self.has_residual): + mC = cute.make_tensor( + c_tensor.iterator, cute.select(c_tensor.layout, mode=[3, 2, 1, 0, 4]) + ) + mC = cute.group_modes(mC, begin=0, end=4) + tma_atom_c, tma_tensor_c = cpasync.make_im2col_tma_atom( + cpasync.CopyBulkTensorIm2ColG2SOp(), + mC, + epi_smem_layout, + self.epi_tile, + lower_corner_whd=(0, 0, 0), + upper_corner_whd=(0, 0, 0), + lower_padding_whd=(0, 0, 0), + upper_padding_whd=(0, 0, 0), + stride_whd=(1, 1, 1), + lower_srt=(0, 0, 0), + stride_srt=(1, 1, 1), + ) + tma_tensor_c = cute.coalesce(tma_tensor_c, target_profile=(1, 1)) + else: + tma_atom_c, tma_tensor_c = None, None + + def add_dummy_batch_dimension(tensor): + new_layout = cute.append(tensor.layout, cute.make_layout(1)) + return cute.make_tensor(tensor.iterator, new_layout) + + tma_tensor_a_preferred = add_dummy_batch_dimension(tma_tensor_a_preferred) + tma_tensor_a_fallback = add_dummy_batch_dimension(tma_tensor_a_fallback) + tma_tensor_b_preferred = add_dummy_batch_dimension(tma_tensor_b_preferred) + tma_tensor_b_fallback = add_dummy_batch_dimension(tma_tensor_b_fallback) + tma_tensor_d = add_dummy_batch_dimension(tma_tensor_d) + if cutlass.const_expr(self.has_residual): + tma_tensor_c = add_dummy_batch_dimension(tma_tensor_c) + + return ( + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + tma_atom_b_preferred, + tma_tensor_b_preferred, + tma_atom_b_fallback, + tma_tensor_b_fallback, + tma_atom_d, + tma_tensor_d, + tma_atom_c, + tma_tensor_c, + ) + + def _setup_attributes(self): + """Set up configurations that are dependent on convolution inputs.""" + # Compute mma instruction shapes + # (MMA_Tile_Shape_M, MMA_Tile_Shape_N, MMA_Inst_Shape_K) + self.mma_inst_shape_mn = ( + self.mma_tiler[0], + self.mma_tiler[1], + ) + # (CTA_Tile_Shape_M, Round_Up(MMA_Tile_Shape_N, 128), MMA_Inst_Shape_K) + self.mma_inst_shape_mn_sfb = ( + self.mma_inst_shape_mn[0] // (2 if self.use_2cta_instrs else 1), + cute.round_up(self.mma_inst_shape_mn[1], 128), + ) + + tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.a_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + self.mma_inst_shape_mn, + ) + + tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.a_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + cute.nvgpu.tcgen05.CtaGroup.ONE, + self.mma_inst_shape_mn_sfb, + ) + + # Compute mma/cluster/tile shapes + mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2]) + # cta_tile_k is a compile-time choice (64, 128 or 256 for FP4). It sets + # how many MMA-instruction K blocks a K tile spans. + mma_inst_tile_k = self.cta_tile_k // mma_inst_shape_k + # Expose for overlapping-accum SF column accounting (see below). + self.mma_inst_tile_k = mma_inst_tile_k + self.mma_tiler = ( + self.mma_tiler[0], + self.mma_tiler[1], + (mma_inst_shape_k * mma_inst_tile_k,), + ) + # Scale factors tile K on their own cadence: an SF atom cannot be staged + # in part, so a narrower A/B K tile still stages a whole one and the A/B + # tiles sharing it each read their own part. Only MXFP8 gets there, with a + # 128-channel atom against a 64 tile; NVFP4's atom is 64, at or below + # every legal tile, so its ratio is 1 and every SF path keeps the A/B one. + self.sf_tile_k = sf_k_tile_channels(self.cta_tile_k, self.sf_vec_size) + self.ab_k_tiles_per_sf_k_tile = self.sf_tile_k // self.cta_tile_k + self.sf_mma_inst_tile_k = self.sf_tile_k // mma_inst_shape_k + # Flat-K tiler the SF layouts (smem, tmem, TMA box) are built from. + self.flat_mma_tiler_sf = ( + self.mma_tiler[0], + self.mma_tiler[1], + self.sf_tile_k, + ) + self.mma_tiler_sfb = ( + self.mma_inst_shape_mn_sfb[0], + self.mma_inst_shape_mn_sfb[1], + self.sf_tile_k, + ) + # SFA 4-byte cp.async per SF K tile: each carries 4 SF blocks, and an SF + # K tile is a whole number of atoms at every legal tile_k, so the count is + # at least one and the power of two the group rotation below needs. + self.num_sfa_cpasync = self.sf_tile_k // (self.sf_vec_size * 4) + self.cta_tile_shape_mnk = ( + self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler[1], + self.mma_tiler[2], + ) + # CTA-level SFB tile shape (used for SFB TMEM column accounting in + # overlapping-accum). N-mode rounds MMA_N up to 128 like SFB MMA. + self.cta_tile_shape_mnk_sfb = ( + self.mma_tiler_sfb[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler_sfb[1], + self.sf_tile_k, + ) + + # Compute cluster layout + self.cluster_layout_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma.thr_id.shape,), + ) + self.cluster_layout_sfb_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma_sfb.thr_id.shape,), + ) + + # Compute number of multicast CTAs for A/B + self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2]) + self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1]) + self.is_a_mcast = self.num_mcast_ctas_a > 1 + self.is_b_mcast = self.num_mcast_ctas_b > 1 + + # Compute epilogue subtile + self.epi_tile = utils.sm100.compute_epilogue_tile_shape( + self.cta_tile_shape_mnk, + self.use_2cta_instrs, + self.d_layout, + self.d_dtype, + ) + # N-extent of one epilogue subtile (used by overlapping-accum early release). + self.epi_tile_n = cute.size(self.epi_tile[1]) + + self.smem_capacity = get_smem_capacity_in_bytes() + + # Setup A/B/D stage count in shared memory and ACC stage count in tensor memory + self.num_acc_stage, self.num_ab_stage, self.num_d_stage = self._compute_stages( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.b_dtype, + self.epi_tile, + self.d_dtype, + self.d_layout, + self.sf_dtype, + self.sf_vec_size, + self.flat_mma_tiler_sf, + self.smem_capacity, + self.occupancy, + self.has_residual, + ) + + # Overlapping-accum: when only one accumulator stage fits in TMEM + # (MMA_N == 256), squeeze a second logical acc buffer into the columns + # otherwise reserved for SFA/SFB so the math mainloop of the next tile + # can overlap the epilogue of the current tile. + self.overlapping_accum = self.num_acc_stage == 1 + # SF columns follow the SF K cadence: TMEM holds whole chunks. + self.num_sfa_tmem_cols = ( + self.cta_tile_shape_mnk[0] // 32 + ) * self.sf_mma_inst_tile_k + self.num_sfb_tmem_cols = ( + self.cta_tile_shape_mnk_sfb[1] // 32 + ) * self.sf_mma_inst_tile_k + self.num_sf_tmem_cols = self.num_sfa_tmem_cols + self.num_sfb_tmem_cols + # Reverse-subtile index (raw loop counter) at which acc can release + # early: the number of whole epilogue subtiles the SFA/SFB-aliased + # columns span. Once the reverse walk has drained that many subtiles the + # shared columns are fully read out and the next tile's MMA may reuse + # them. When the shared columns fit within one subtile (num_sf_tmem_cols + # <= epi_tile_n) this is 0, i.e. release right after the first subtile. + self.iter_acc_early_release_in_epilogue = ( + self.num_sf_tmem_cols // self.epi_tile_n + ) + + # Setup CLC stage (single-stage CLC pipeline) + self.num_clc_stage = 1 + assert self.num_clc_stage == 1, "Only single-stage CLC pipeline is supported" + + # Compute A/B/D shared memory layout + ( + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.d_smem_layout_staged, + ) = self._make_smem_layouts( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.b_dtype, + self.d_dtype, + self.d_layout, + self.epi_tile, + self.num_ab_stage, + self.num_d_stage, + ) + self.sfa_smem_layout_staged, self.sfb_smem_layout_staged = ( + self._make_sf_smem_layouts( + tiled_mma, + self.flat_mma_tiler_sf, + self.sf_vec_size, + self.num_ab_stage, + ) + ) + + # Compute the number of tensor memory allocation columns. + # For overlapping-accum the second acc buffer is strided into the + # SFA/SFB columns, so the inherited helper (which builds a contiguous + # 2-stage fake) would under-count. Derive the column count from the + # same fake tensor used at the injection sites instead. + if cutlass.const_expr(self.overlapping_accum): + self.num_tmem_alloc_cols = get_num_tmem_alloc_cols( + self._make_acc_fake_tensor(tiled_mma, self.mma_tiler), + arch=self.arch, + ) + else: + self.num_tmem_alloc_cols = self._compute_num_tmem_alloc_cols( + tiled_mma, self.mma_tiler, self.num_acc_stage, self.arch + ) + + # Compute preferred cluster layout for dual-cluster scheduling + self.preferred_cluster_layout_vmnk = cute.tiled_divide( + cute.make_layout((*self.preferred_cluster_shape_mn, 1)), + (tiled_mma.thr_id.shape,), + ) + self.num_preferred_mcast_ctas_a = cute.size( + self.preferred_cluster_layout_vmnk.shape[2] + ) + self.num_preferred_mcast_ctas_b = cute.size( + self.preferred_cluster_layout_vmnk.shape[1] + ) + self.is_preferred_a_mcast = self.num_preferred_mcast_ctas_a > 1 + self.is_preferred_b_mcast = self.num_preferred_mcast_ctas_b > 1 + + # Fallback cluster layout was already computed as cluster_layout_vmnk above + self.fallback_cluster_layout_vmnk = self.cluster_layout_vmnk + self.num_fallback_mcast_ctas_a = self.num_mcast_ctas_a + self.num_fallback_mcast_ctas_b = self.num_mcast_ctas_b + self.is_fallback_a_mcast = self.is_a_mcast + self.is_fallback_b_mcast = self.is_b_mcast + + # SFB cluster layouts (SFB always uses 1CTA group, thr_id shape = 1) + self.preferred_cluster_layout_sfb_vmnk = cute.tiled_divide( + cute.make_layout((*self.preferred_cluster_shape_mn, 1)), + (tiled_mma_sfb.thr_id.shape,), + ) + self.fallback_cluster_layout_sfb_vmnk = self.cluster_layout_sfb_vmnk + + @cute.jit + def __call__( + self, + a_tensor: cute.Tensor, + b_tensor: cute.Tensor, + d_tensor: cute.Tensor, + sfa_tensor: cute.Tensor, + sfb_tensor: cute.Tensor, + alpha: cutlass.Float32, + epilogue_op: cutlass.Constexpr = lambda x: x, + sfd_tensor: Optional[cute.Tensor] = None, + norm_const: cutlass.Float32 = 1.0, + bias_tensor: Optional[cute.Tensor] = None, + c_tensor: Optional[cute.Tensor] = None, + beta: cutlass.Constexpr = 0.0, + # Conv pad/stride/dilation as cute.compile entry scalars. The caller + # passes boxed cutlass.Int32(...) so they lower to runtime SSA: one + # cubin then serves any pad/stride/dil because the SAME scalars feed + # BOTH the im2col A descriptor corner (host-side, in _setup_conv_tma) + # AND the device per-tile coord math. Boxing matters: a raw Python int + # here would embed its value in the mangled function name, defeating + # reuse and forcing recompilation per config. + rt_upper_pad_d: cutlass.Int32 = None, + rt_upper_pad_h: cutlass.Int32 = None, + rt_upper_pad_w: cutlass.Int32 = None, + rt_lower_pad_d: cutlass.Int32 = None, + rt_lower_pad_h: cutlass.Int32 = None, + rt_lower_pad_w: cutlass.Int32 = None, + rt_stride_d: cutlass.Int32 = None, + rt_stride_h: cutlass.Int32 = None, + rt_stride_w: cutlass.Int32 = None, + rt_dil_d: cutlass.Int32 = None, + rt_dil_h: cutlass.Int32 = None, + rt_dil_w: cutlass.Int32 = None, + stream: cuda.CUstream = None, + ): + """Execute the persistent convolution operation with dynamic preferred cluster scheduling. + + :param a_tensor: Input tensor A - (N, D, H, W, C) layout + :param b_tensor: Filter tensor B - (K, T, R, S, C) layout + :param d_tensor: Output tensor D - (N, Z, P, Q, K) layout + :param sfa_tensor: Block scale factor tensor for A + :param sfb_tensor: Block scale factor tensor for B + :param alpha: FP32 runtime scalar; applied to the accumulator in FP32 + before output quantization + :param epilogue_op: Optional elementwise lambda function to apply to the output tensor + :param sfd_tensor: Output scale factor tensor (NVFP4 SFD); pass None to skip SFD generation + :param norm_const: FP32 runtime scalar (default 1.0); per-tensor amax-derived + global FP4 scale, multiplied into the SFD encode step. Only used when + sfd_tensor is provided. + :param bias_tensor: Optional length-K (output channel) device tensor; a + per-output-channel bias added to the accumulator in FP32 as + D = epilogue_op(alpha * acc + bias). Pass None to skip bias. + :param c_tensor: Optional (N, Z, P, Q, K) device tensor with the + same shape/layout as the output D; a per-element residual added to + the accumulator in FP32 as D = epilogue_op(alpha * acc + bias + + beta * residual). Loaded through the epilogue via a TMA im2col G2S + copy (the inverse of the output store) into shared memory, then read + to registers. Required when beta != 0. + :param beta: compile-time residual scaling constant. beta == 0 removes + the entire residual path (smem, pipeline, TMA load, add) via const + DCE; beta != 0 scales the residual as D = alpha * acc + bias + + beta * residual and requires c_tensor. A distinct beta value + yields a distinct cubin. + :param stream: CUDA stream for asynchronous execution; defaults to None + """ + self._setup_conv_input_attrs(a_tensor, b_tensor, d_tensor) + self.sf_dtype: Type[cutlass.Numeric] = sfa_tensor.element_type + # SFD shares the input sf_dtype (E4M3 for NVFP4, E8M0 for MXFP8). + self.sfd_dtype: Type[cutlass.Numeric] = self.sf_dtype + # SFD opt-in: caller supplies sfd_tensor. SFD is valid only with a narrow + # output (FP4 or FP8 E4M3). norm_const is a plain scalar (default 1.0) + # folded into the SFD scale-up. + self.gen_sfd: bool = sfd_tensor is not None + # Bias opt-in: caller supplies a length-K per-output-channel tensor. + self.has_bias: bool = bias_tensor is not None + # Residual opt-in via compile-time beta: beta == 0 removes the whole + # residual path through const DCE; beta != 0 scales the per-element + # residual as alpha*acc + bias + beta*residual and needs c_tensor. + self.beta: cutlass.Constexpr = beta + self.has_residual: bool = beta != 0.0 + if cutlass.const_expr(self.has_residual and c_tensor is None): + raise ValueError("beta != 0 requires a c_tensor") + # Finalize the CTA thread count now that has_residual is known: the + # dedicated residual G2S warp only exists when beta != 0, so beta == 0 + # launches with no extra idle warp. + if cutlass.const_expr(self.has_residual): + self.threads_per_cta = 32 * (len(self.base_warp_ids) + 1) + else: + self.threads_per_cta = 32 * len(self.base_warp_ids) + if cutlass.const_expr( + self.gen_sfd + and self.d_dtype not in (cutlass.Float4E2M1FN, cutlass.Float8E4M3FN) + ): + raise ValueError( + "SFD output is only supported for Float4E2M1FN (NVFP4) or " + f"Float8E4M3FN (MXFP8) output; got d_dtype={self.d_dtype}" + ) + # sf_dtype is derived from SFA alone but also drives the SFB smem layout + # and barrier transaction bytes; SFB/SFD must match it or scales get + # misinterpreted and barrier byte counts diverge. + if cutlass.const_expr(sfb_tensor.element_type is not self.sf_dtype): + raise ValueError( + f"SFB dtype ({sfb_tensor.element_type}) must match SFA/" + f"sf_dtype ({self.sf_dtype})" + ) + if cutlass.const_expr( + self.gen_sfd and sfd_tensor.element_type is not self.sfd_dtype + ): + raise ValueError( + f"SFD dtype ({sfd_tensor.element_type}) must match " + f"sf_dtype ({self.sfd_dtype})" + ) + # sBias is allocated with d_dtype but the staging cp.async width is + # taken from the bias tensor's dtype, so a mismatch overruns the smem + # buffer. Require bias to match the output dtype. + if cutlass.const_expr( + self.has_bias and bias_tensor.element_type is not self.d_dtype + ): + raise ValueError( + f"bias dtype ({bias_tensor.element_type}) must match output " + f"d_dtype ({self.d_dtype})" + ) + self._setup_attributes() + + # Pad/stride/dilation operands that BOTH the im2col A descriptor corner + # and the device per-tile coord math consume. They are the runtime + # Int32 scalars passed at the cute.compile entry, so one cubin serves + # any config. Threading the SAME values into the descriptor and the + # kernel is required: runtime-ing only one side would let A's load + # coords (descriptor) desync from SFA's cp.async coords (device), + # zeroing the accumulator on any non-default config. + upper_pad_op = (rt_upper_pad_d, rt_upper_pad_h, rt_upper_pad_w) + lower_pad_op = (rt_lower_pad_d, rt_lower_pad_h, rt_lower_pad_w) + stride_op = (rt_stride_d, rt_stride_h, rt_stride_w) + dil_op = (rt_dil_d, rt_dil_h, rt_dil_w) + + tiled_mma = sm100_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.a_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + self.cta_group, + self.mma_tiler[:2], + ) + + ( + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + tma_atom_b_preferred, + tma_tensor_b_preferred, + tma_atom_b_fallback, + tma_tensor_b_fallback, + tma_atom_d, + tma_tensor_d, + tma_atom_c, + tma_tensor_c, + ) = self._setup_conv_tma( + a_tensor, + b_tensor, + d_tensor, + c_tensor, + tiled_mma, + upper_pad_op, + lower_pad_op, + stride_op, + dil_op, + ) + + # SFB tiled_mma (always 1CTA group for SFB) + tiled_mma_sfb = sm100_utils.make_blockscaled_trivial_tiled_mma( + self.a_dtype, + self.a_dtype, + self.a_major_mode, + self.b_major_mode, + self.sf_dtype, + self.sf_vec_size, + cute.nvgpu.tcgen05.CtaGroup.ONE, + self.mma_inst_shape_mn_sfb, + ) + atom_thr_size = cute.size(tiled_mma.thr_id.shape) + + # SFB tensor reshape: match B's GEMM view (K, C*S*R*T) for correct TMA tiling. + # B filter is (K, T, R, S, C) -> reorder to (K, C, S, R, T) -> group to (K, C*S*R*T). + # This makes SFB N dim = K (GEMM N) and SFB K dim = C*S*R*T/16 (GEMM K / sf_vec). + # Use a flat integer K extent for tile_atom_to_shape_SF, NOT the hierarchical + # shape from group_modes: group_modes produces (K, (C, S, R, T)) with nested + # strides, which tile_atom_to_shape_SF interprets incorrectly vs the flat + # swizzled SFB storage. + # + # The N extent MUST be the dynamic b.shape[0], not a trace-time-static int. + # For cta_n==192 the SFB layout is reshaped into overlapping 256-wide windows + # ((2,2),y) with y=ceil_div(N_sf,4) below. A static N collapses y to a literal + # 1, and CuTe coalesces the size-1 sub-mode away, flattening the window RestN + # from (2,?) to 2. The scheduler still emits ceil(N/192) cta-tiles, so for a + # partial last tile slice_n runs OOB on the flattened RestN and the TMA load + # clamps to zeros -> wrong (zero) SFB scaling on that tile. A dynamic N keeps y + # symbolic, blocks the coalesce, and slice_n decomposes to an in-bounds coord. + # C*T*R*S = GEMM-K; T/R/S come from the dynamic filter tensor (runtime). The + # channel extent is the span SFB was allocated at -- C rounded up to whole SF + # K tiles -- on two counts. The N-mode rest stride tile_atom_to_shape_SF + # derives is the tensor's whole atom count, so a short extent walks the + # output-channel axis at the wrong pitch. And the rounding puts every filter + # position on a tile boundary, which keeps a K tile inside one position and + # lets the consumer's flat K index address the tile directly. The span comes + # from the runtime channel count, so one cubin serves every C. + sfb_c_extent = cute.ceil_div(b_tensor.shape[4], self.sf_tile_k) * self.sf_tile_k + b_shape_for_sfb = ( + b_tensor.shape[0], + sfb_c_extent * b_tensor.shape[1] * b_tensor.shape[2] * b_tensor.shape[3], + 1, + ) + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF( + b_shape_for_sfb, self.sf_vec_size + ) + sfb_tensor = cute.make_tensor(sfb_tensor.iterator, sfb_layout) + + # SFB TMA setup (dual atoms for preferred + fallback cluster shapes) + sfb_smem_layout = cute.slice_( + self.sfb_smem_layout_staged, (None, None, None, 0) + ) + + def _make_sfb_tma(cluster_shape_mn, cluster_layout_sfb_vmnk): + sfb_op = sm100_utils.cluster_shape_to_tma_atom_SFB( + cluster_shape_mn, tiled_mma.thr_id + ) + atom, tensor = cute.nvgpu.make_tiled_tma_atom_B( + sfb_op, + sfb_tensor, + sfb_smem_layout, + self.mma_tiler_sfb, + tiled_mma_sfb, + cluster_layout_sfb_vmnk.shape, + internal_type=cutlass.Int16, + ) + # Special-case for N=192: align the SFB scale-factor layout for Tensor Memory packing + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + x = tensor.stride[0][1] + y = cute.ceil_div(tensor.shape[0][1], 4) + new_shape = ( + (tensor.shape[0][0], ((2, 2), y)), + tensor.shape[1], + tensor.shape[2], + ) + x_times_3 = 3 * x + new_stride = ( + (tensor.stride[0][0], ((x, x), x_times_3)), + tensor.stride[1], + tensor.stride[2], + ) + tensor = cute.make_tensor( + tensor.iterator, + cute.make_layout(new_shape, stride=new_stride), + ) + return atom, tensor + + tma_atom_sfb_preferred, tma_tensor_sfb_preferred = _make_sfb_tma( + self.preferred_cluster_shape_mn, self.preferred_cluster_layout_sfb_vmnk + ) + tma_atom_sfb_fallback, tma_tensor_sfb_fallback = _make_sfb_tma( + self.fallback_cluster_shape_mn, self.fallback_cluster_layout_sfb_vmnk + ) + + # Compute A, B, SFB copy sizes for separate TMA pipelines. + # SFB has its own pipeline because B and SFB use different mcast + # masks (cluster_layout_vmnk vs cluster_layout_sfb_vmnk), so bundling + # them into one pipeline can stall TX accumulation in cluster > 1. + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0)) + b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) + a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout) + b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout) + sfb_copy_size = cute.size_in_bytes(self.sf_dtype, sfb_smem_layout) + self.num_tma_load_a_bytes = a_copy_size * atom_thr_size + self.num_tma_load_b_bytes = b_copy_size * atom_thr_size + self.num_tma_load_sfb_bytes = sfb_copy_size * atom_thr_size + + # Residual reuses the D epilogue smem layout (same shape/dtype as the + # output); one epilogue subtile's worth of bytes per TMA load. + c_smem_layout = cute.slice_(self.d_smem_layout_staged, (None, None, 0)) + self.num_tma_load_c_bytes = cute.size_in_bytes(self.d_dtype, c_smem_layout) + + self.buffer_align_bytes = 1024 + + # Define shared storage for kernel + @cute.struct + class SharedStorage: + # Shared TMA pipeline barriers for A, B and SFB + ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage * 2] + # CPASYNC SFA pipeline barriers + sfa_full_mbar_ptr: cute.struct.MemRange[ + cutlass.Int64, self.num_ab_stage * 2 + ] + acc_full_mbar_ptr: cute.struct.MemRange[ + cutlass.Int64, self.num_acc_stage * 2 + ] + tmem_dealloc_mbar_ptr: cutlass.Int64 + tmem_holding_buf: cutlass.Int32 + # CLC pipeline barriers and response buffer + clc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_clc_stage * 2] + clc_response: cute.struct.Align[ + cute.struct.MemRange[ + cutlass.Int32, self.num_clc_response_bytes // 4 * self.num_clc_stage + ], + 16, + ] + # (EPI_TILE_M, EPI_TILE_N, STAGE) + sD: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, + cute.cosize(self.d_smem_layout_staged.outer), + ], + self.buffer_align_bytes, + ] + # (cta_tile_n,) staged bias row for the shared-memory broadcast. + # One value per output channel in the CTA's N-tile; sized to zero + # when no bias is supplied so it costs no smem. + sBias: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, + self.cta_tile_shape_mnk[1] if self.has_bias else 0, + ], + self.buffer_align_bytes, + ] + # (EPI_TILE_M, EPI_TILE_N, STAGE) residual tile, same layout as sD. + # Sized to zero when no residual is supplied so it costs no smem. + sC: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, + cute.cosize(self.d_smem_layout_staged.outer) + if self.has_residual + else 0, + ], + self.buffer_align_bytes, + ] + # Residual TMA load pipeline barriers (full/empty per stage). + c_full_mbar_ptr: cute.struct.MemRange[ + cutlass.Int64, self.num_d_stage * 2 if self.has_residual else 0 + ] + # (MMA, MMA_M, MMA_K, STAGE) + sA: cute.struct.Align[ + cute.struct.MemRange[ + self.a_dtype, cute.cosize(self.a_smem_layout_staged.outer) + ], + self.buffer_align_bytes, + ] + # (MMA, MMA_N, MMA_K, STAGE) + sB: cute.struct.Align[ + cute.struct.MemRange[ + self.b_dtype, cute.cosize(self.b_smem_layout_staged.outer) + ], + self.buffer_align_bytes, + ] + # (MMA, MMA_M, MMA_K, STAGE) + sSFA: cute.struct.Align[ + cute.struct.MemRange[ + self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + # (MMA, MMA_N, MMA_K, STAGE) + sSFB: cute.struct.Align[ + cute.struct.MemRange[ + self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + + self.shared_storage = SharedStorage + + # Compute grid size and scheduler params for both cluster shapes + self.fallback_tile_sched_params, _ = self._compute_grid( + tma_tensor_d, + self.cta_tile_shape_mnk, + self.fallback_cluster_shape_mn, + self.swizzle_size, + self.raster_along, + ) + self.preferred_tile_sched_params, preferred_grid = self._compute_grid( + tma_tensor_d, + self.cta_tile_shape_mnk, + self.preferred_cluster_shape_mn, + self.swizzle_size, + self.raster_along, + ) + # Extract conv spatial dims from input/output tensors + # a_tensor: (N, D, H, W, C), d_tensor: (N, Z, P, Q, K) + conv_N = cute.size(a_tensor, mode=[0]) + conv_D = cute.size(a_tensor, mode=[1]) + conv_H = cute.size(a_tensor, mode=[2]) + conv_W = cute.size(a_tensor, mode=[3]) + conv_C = cute.size(a_tensor, mode=[4]) + conv_Z = cute.size(d_tensor, mode=[1]) + conv_P = cute.size(d_tensor, mode=[2]) + conv_Q = cute.size(d_tensor, mode=[3]) + # Filter T/R/S from the dynamic-layout filter tensor b_tensor (K, T, R, S, C): + # runtime Int32 so one cubin serves any T/R/S. Fed to the device + # kernel's conv_T/R/S params. + conv_T, conv_R, conv_S = b_tensor.shape[1], b_tensor.shape[2], b_tensor.shape[3] + # stride/dil/pad come from the SAME operand tuples that built the A + # descriptor corner above, so the device's SFA cp.async coords stay in + # lockstep with A's im2col load coords under any runtime config. + stride_d, stride_h, stride_w = stride_op + dil_d, dil_h, dil_w = dil_op + # SFA cp.async path uses the leading (lower) pad to mirror what the + # im2col TMA A descriptor subtracts when mapping output->input coords: + # d_in = z*stride - lower_pad + t*dil + # Using upper_padding here desyncs SFA from A on asymmetric pads, + # making A's nonzero positions multiply SFA's zero positions -> acc=0. + pad_d, pad_h, pad_w = lower_pad_op + K_gemm_tile = self.mma_tiler[2][0] # mma_inst_shape_k * mma_inst_tile_k + + # Build mSFD tensor with the (M, K, 1) profile of mD, where M is the + # flat NZPQ output-pixel axis and the K dimension is expressed as + # (sfd_vec_size, sf_k_padded) with strides (0, 1) so that sfd_vec_size + # consecutive K elements share one physical scale factor. Storage layout + # takes SFA's K-contig-in-M form, (mn, sf_k, 1) with mn-stride = sf_k, + # except that sf_k is rounded up here: it counts Kout, which no gate + # constrains, so the 4-byte transfers walking that stride need the tail. + if cutlass.const_expr(self.gen_sfd): + # Source N/Z/P/Q/K from the dynamic output tensor (runtime Int) so + # one cubin serves any output geometry. The SFD store is a scalar + # per-element STG with a static (unrolled) element count, so only the + # per-thread tile is static; the global extent/strides can be runtime. + sfd_N = conv_N + sfd_Z = conv_Z + sfd_P = conv_P + sfd_Q = conv_Q + conv_K = cute.size(d_tensor, mode=[4]) + sf_k = (conv_K + self.sfd_vec_size - 1) // self.sfd_vec_size + sf_k_padded = ((sf_k + 3) // 4) * 4 # pad sf_k to multiple of 4 + stride_mn = sf_k_padded + # The M axis is a single flat extent N*Z*P*Q with a uniform row + # stride, not a hierarchical (Q,P,Z,N) tuple. local_tile splits the + # leading M mode by the CTA tile, and a runtime-valued hierarchical + # mode does not divide correctly there: the store then lands on wrong + # rows whenever N>1 (or any spatial extent grows past one tile), + # leaving some scale-factor rows unwritten. A flat runtime extent + # tiles cleanly because the SFD gmem buffer is dense row-major NZPQ + # and a plain STG needs no hierarchical structure. + sfd_M = sfd_N * sfd_Z * sfd_P * sfd_Q + mSFD_layout = cute.make_layout( + ( + sfd_M, + (self.sfd_vec_size, sf_k_padded), + 1, + ), + stride=( + stride_mn, + (0, 1), + sfd_M * stride_mn, + ), + ) + mSFD_mnl = cute.make_tensor(sfd_tensor.iterator, mSFD_layout) + else: + mSFD_mnl = None + + # Build mBias tensor: a per-output-channel bias broadcast to the same + # ((Q,P,Z,N), K, 1) = (M, N, L) profile as mD. The bias varies along the + # output channel K (GEMM-N, real stride 1) and broadcasts across every + # spatial output position (GEMM-M = Q,P,Z,N modes carry stride 0), so the + # epilogue reads it through the same partition_C chain as the accumulator + # and each thread lands on its own N-column bias scalar. + if cutlass.const_expr(bias_tensor is not None): + # The M (spatial) modes are stride-0 broadcast: every output row + # reads the same bias[k], so their extent never enters an address and + # can be a compile-time value. Using a static M extent keeps the + # partitioned smem-read layout static (needed for a vectorized smem load) + # while the N (K/channel) axis carries a runtime stride so one cubin + # serves any channel count. The M extent only needs to cover the CTA + # tile; the real output-M size is supplied at the store site through + # the runtime tile coordinate, not through this layout. + mBias_layout = cute.make_layout( + ( + (self.cta_tile_shape_mnk[0], 1, 1, 1), + cute.size(d_tensor, mode=[4]), + 1, + ), + stride=((0, 0, 0, 0), 1, 0), + ) + mBias_mnl = cute.make_tensor(bias_tensor.iterator, mBias_layout) + else: + mBias_mnl = None + + # Launch the megakernel synchronously (dispatcher selects preferred vs fallback) + self.kernel( + tiled_mma, + tiled_mma_sfb, + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + tma_atom_b_preferred, + tma_tensor_b_preferred, + tma_atom_b_fallback, + tma_tensor_b_fallback, + sfa_tensor, + tma_atom_sfb_preferred, + tma_tensor_sfb_preferred, + tma_atom_sfb_fallback, + tma_tensor_sfb_fallback, + tma_atom_d, + tma_tensor_d, + tma_atom_c, + tma_tensor_c, + mSFD_mnl, + alpha, + norm_const, + mBias_mnl, + self.preferred_cluster_layout_vmnk, + self.fallback_cluster_layout_vmnk, + self.preferred_cluster_layout_sfb_vmnk, + self.fallback_cluster_layout_sfb_vmnk, + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.sfa_smem_layout_staged, + self.sfb_smem_layout_staged, + self.d_smem_layout_staged, + self.epi_tile, + self.preferred_tile_sched_params, + self.fallback_tile_sched_params, + epilogue_op, + conv_T, + conv_R, + conv_S, + stride_d, + stride_h, + stride_w, + dil_d, + dil_h, + dil_w, + pad_d, + pad_h, + pad_w, + conv_D, + conv_H, + conv_W, + conv_Z, + conv_P, + conv_Q, + conv_N, + conv_C, + K_gemm_tile, + ).launch( + grid=preferred_grid, + block=[self.threads_per_cta, 1, 1], + cluster=(*self.preferred_cluster_shape_mn, 1), + fallback_cluster=(*self.fallback_cluster_shape_mn, 1), + stream=stream, + smem_merge_branch_allocs=True, + ) + return + + # GPU device kernel - megakernel dispatcher for preferred/fallback cluster + @cute.kernel + def kernel( + self, + tiled_mma: cute.TiledMma, + tiled_mma_sfb: cute.TiledMma, + tma_atom_a_preferred: cute.CopyAtom, + mA_mkl_preferred: cute.Tensor, + tma_atom_a_fallback: cute.CopyAtom, + mA_mkl_fallback: cute.Tensor, + tma_atom_b_preferred: cute.CopyAtom, + mB_nkl_preferred: cute.Tensor, + tma_atom_b_fallback: cute.CopyAtom, + mB_nkl_fallback: cute.Tensor, + mSFA_mkl: cute.Tensor, + tma_atom_sfb_preferred: cute.CopyAtom, + mSFB_nkl_preferred: cute.Tensor, + tma_atom_sfb_fallback: cute.CopyAtom, + mSFB_nkl_fallback: cute.Tensor, + tma_atom_d: cute.CopyAtom, + mD_mnl: cute.Tensor, + tma_atom_c: Optional[cute.CopyAtom], + mC_mnl: Optional[cute.Tensor], + mSFD_mnl: Optional[cute.Tensor], + alpha: cutlass.Float32, + norm_const: cutlass.Float32, + mBias_mnl: Optional[cute.Tensor], + preferred_cluster_layout_vmnk: cute.Layout, + fallback_cluster_layout_vmnk: cute.Layout, + preferred_cluster_layout_sfb_vmnk: cute.Layout, + fallback_cluster_layout_sfb_vmnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + sfa_smem_layout_staged: cute.Layout, + sfb_smem_layout_staged: cute.Layout, + d_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout], + epi_tile: cute.Tile, + preferred_tile_sched_params: ClcDynamicPersistentTileSchedulerParams, + fallback_tile_sched_params: ClcDynamicPersistentTileSchedulerParams, + epilogue_op: cutlass.Constexpr, + # Conv parameters for SFA im2col coordinate mapping. conv_T/R/S are + # runtime Int32 so one cubin serves any filter size; they only feed + # register comparisons in the K-loop (S->R->T wrap), adding no IDIV. + conv_T: cutlass.Int32, + conv_R: cutlass.Int32, + conv_S: cutlass.Int32, + # pad/stride/dilation are runtime Int32 so one cubin serves any conv + # config; they feed the per-tile base-coord arithmetic (z*stride - pad) + # and the K-loop dilation step, swapping immediates for registers with + # no new IDIV. + stride_d: cutlass.Int32, + stride_h: cutlass.Int32, + stride_w: cutlass.Int32, + dil_d: cutlass.Int32, + dil_h: cutlass.Int32, + dil_w: cutlass.Int32, + pad_d: cutlass.Int32, + pad_h: cutlass.Int32, + pad_w: cutlass.Int32, + # Coordinate geometry: runtime Int32 so the per-tile NDHW decomposition + # (m_global -> n,z,p,q) lowers to real IDIV and one cubin serves any + # N/D/H/W/Z/P/Q. + input_D: cutlass.Int32, + input_H: cutlass.Int32, + input_W: cutlass.Int32, + output_Z: cutlass.Int32, + output_P: cutlass.Int32, + output_Q: cutlass.Int32, + input_N: cutlass.Int32, + input_C: cutlass.Int32, + K_gemm_tile: cutlass.Constexpr, + ): + """ + GPU device kernel entry point: dispatches to preferred or fallback cluster path. + """ + # Determine if this CTA is launched in a preferred-shape cluster + cbdim_x, cbdim_y, cbdim_z = cute.arch.block_in_cluster_dim() + is_preferred_cluster = ( + cbdim_x == self.preferred_cluster_shape_mn[0] + and cbdim_y == self.preferred_cluster_shape_mn[1] + and cbdim_z == 1 + ) + + # Megakernel: only one branch executes per launch. + # smem_merge_branch_allocs=True at launch enables shared memory reuse between two paths. + if is_preferred_cluster: + self.cluster_specific_kernel( + tiled_mma, + tiled_mma_sfb, + tma_atom_a_preferred, + mA_mkl_preferred, + tma_atom_b_preferred, + mB_nkl_preferred, + mSFA_mkl, + tma_atom_sfb_preferred, + mSFB_nkl_preferred, + tma_atom_d, + mD_mnl, + tma_atom_c, + mC_mnl, + mSFD_mnl, + alpha, + norm_const, + mBias_mnl, + preferred_cluster_layout_vmnk, + preferred_cluster_layout_sfb_vmnk, + a_smem_layout_staged, + b_smem_layout_staged, + sfa_smem_layout_staged, + sfb_smem_layout_staged, + d_smem_layout_staged, + epi_tile, + preferred_tile_sched_params, + epilogue_op, + self.num_preferred_mcast_ctas_a + self.num_preferred_mcast_ctas_b - 1, + self.is_preferred_a_mcast, + self.is_preferred_b_mcast, + self.preferred_cluster_shape_mn, + conv_T, + conv_R, + conv_S, + stride_d, + stride_h, + stride_w, + dil_d, + dil_h, + dil_w, + pad_d, + pad_h, + pad_w, + input_D, + input_H, + input_W, + output_Z, + output_P, + output_Q, + input_N, + input_C, + K_gemm_tile, + ) + else: + self.cluster_specific_kernel( + tiled_mma, + tiled_mma_sfb, + tma_atom_a_fallback, + mA_mkl_fallback, + tma_atom_b_fallback, + mB_nkl_fallback, + mSFA_mkl, + tma_atom_sfb_fallback, + mSFB_nkl_fallback, + tma_atom_d, + mD_mnl, + tma_atom_c, + mC_mnl, + mSFD_mnl, + alpha, + norm_const, + mBias_mnl, + fallback_cluster_layout_vmnk, + fallback_cluster_layout_sfb_vmnk, + a_smem_layout_staged, + b_smem_layout_staged, + sfa_smem_layout_staged, + sfb_smem_layout_staged, + d_smem_layout_staged, + epi_tile, + fallback_tile_sched_params, + epilogue_op, + self.num_fallback_mcast_ctas_a + self.num_fallback_mcast_ctas_b - 1, + self.is_fallback_a_mcast, + self.is_fallback_b_mcast, + self.fallback_cluster_shape_mn, + conv_T, + conv_R, + conv_S, + stride_d, + stride_h, + stride_w, + dil_d, + dil_h, + dil_w, + pad_d, + pad_h, + pad_w, + input_D, + input_H, + input_W, + output_Z, + output_P, + output_Q, + input_N, + input_C, + K_gemm_tile, + ) + + @cute.jit() + def cluster_specific_kernel( + self, + tiled_mma: cute.TiledMma, + tiled_mma_sfb: cute.TiledMma, + tma_atom_a: cute.CopyAtom, + mA_mkl: cute.Tensor, + tma_atom_b: cute.CopyAtom, + mB_nkl: cute.Tensor, + mSFA_mkl: cute.Tensor, + tma_atom_sfb: cute.CopyAtom, + mSFB_nkl: cute.Tensor, + tma_atom_d: cute.CopyAtom, + mD_mnl: cute.Tensor, + tma_atom_c: Optional[cute.CopyAtom], + mC_mnl: Optional[cute.Tensor], + mSFD_mnl: Optional[cute.Tensor], + alpha: cutlass.Float32, + norm_const: cutlass.Float32, + mBias_mnl: Optional[cute.Tensor], + cluster_layout_vmnk: cute.Layout, + cluster_layout_sfb_vmnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + sfa_smem_layout_staged: cute.Layout, + sfb_smem_layout_staged: cute.Layout, + d_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout], + epi_tile: cute.Tile, + tile_sched_params: ClcDynamicPersistentTileSchedulerParams, + epilogue_op: cutlass.Constexpr, + num_tma_producer: int, + effective_is_a_mcast: bool, + effective_is_b_mcast: bool, + cluster_shape: Tuple[int, int], + # Conv parameters for SFA im2col coordinate mapping. conv_T/R/S are + # runtime Int32 so one cubin serves any filter size; they only feed + # register comparisons in the K-loop (S->R->T wrap), adding no IDIV. + conv_T: cutlass.Int32, + conv_R: cutlass.Int32, + conv_S: cutlass.Int32, + # pad/stride/dilation are runtime Int32 so one cubin serves any conv + # config; they feed the per-tile base-coord arithmetic (z*stride - pad) + # and the K-loop dilation step, swapping immediates for registers with + # no new IDIV. + stride_d: cutlass.Int32, + stride_h: cutlass.Int32, + stride_w: cutlass.Int32, + dil_d: cutlass.Int32, + dil_h: cutlass.Int32, + dil_w: cutlass.Int32, + pad_d: cutlass.Int32, + pad_h: cutlass.Int32, + pad_w: cutlass.Int32, + # Coordinate geometry: runtime Int32 so the per-tile NDHW decomposition + # (m_global -> n,z,p,q) lowers to real IDIV and one cubin serves any + # N/D/H/W/Z/P/Q. + input_D: cutlass.Int32, + input_H: cutlass.Int32, + input_W: cutlass.Int32, + output_Z: cutlass.Int32, + output_P: cutlass.Int32, + output_Q: cutlass.Int32, + input_N: cutlass.Int32, + input_C: cutlass.Int32, + K_gemm_tile: cutlass.Constexpr, + ): + """ + GPU device kernel performing the CLC dynamic persistent convolution computation. + """ + warp_idx = cute.arch.warp_idx() + warp_idx = cute.arch.make_warp_uniform(warp_idx) + + # + # Prefetch tma desc + # + if warp_idx == self.tma_warp_id: + cpasync.prefetch_descriptor(tma_atom_a) + cpasync.prefetch_descriptor(tma_atom_b) + # SFA has no TMA descriptor to prefetch: it is loaded via cp.async. + cpasync.prefetch_descriptor(tma_atom_sfb) + cpasync.prefetch_descriptor(tma_atom_d) + + use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2 + + # + # Setup cta/thread coordinates + # + # Coords inside cluster + bidx, bidy, bidz = cute.arch.block_idx() + mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape) + is_leader_cta = mma_tile_coord_v == 0 + cta_rank_in_cluster = cute.arch.make_warp_uniform( + cute.arch.block_idx_in_cluster() + ) + is_first_cta_in_cluster = cta_rank_in_cluster == 0 + block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord( + cta_rank_in_cluster + ) + block_in_cluster_coord_sfb_vmnk = cluster_layout_sfb_vmnk.get_flat_coord( + cta_rank_in_cluster + ) + # Coord inside cta + tidx, _, _ = cute.arch.thread_idx() + + # + # Alloc and init: a+b full/empty, accumulator full/empty, tensor memory dealloc barrier + # + smem = SmemAllocator() + storage = smem.allocate(self.shared_storage) + + # Pipeline Init: merged pipeline shares one mbarrier between TMA (A+B+SFB) + # and cp.async (SFA) producers, for any cluster shape. + # tx_count covers A+B+SFB TMA bytes; arrive_count = 128 cp.async lanes + + # 1 TMA arrive. The TMA arrive is produced by the leader's multicast + # arrive_and_expect_tx and propagated to every peer barrier by hardware, + # so peer CTAs must not issue an additional arrive. + # For 2CTA (use_2cta_instrs), peer CTA's 128 cpasync threads also + # remote-arrive on leader's sfa_full via DSMEM (shift-by-K pattern) + # to eliminate the peer-sSFA race where leader MMA proceeded before + # peer's cp.async landed. Leader's arrive_count must therefore be + # bumped to 129 + 128 = 257. + # num_tma_producer is provided as kernel parameter (per-cluster) + sfa_cpasync_arrive_count = 257 if use_2cta_instrs else 129 + cpasynd_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, sfa_cpasync_arrive_count + ) + merged_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_tma_producer + ) + merged_tx_count = ( + self.num_tma_load_a_bytes + + self.num_tma_load_b_bytes + + self.num_tma_load_sfb_bytes + ) + ab_sfa_pipeline = PipelineTmaCpAsyncUmma.create( + barrier_storage=storage.ab_full_mbar_ptr.data_ptr(), + num_stages=self.num_ab_stage, + cpasynd_producer_group=cpasynd_producer_group, + consumer_group=merged_consumer_group, + tx_count=merged_tx_count, + cta_layout_vmnk=cluster_layout_vmnk, + defer_sync=True, + ) + + # Pipeline Init: Initialize acd_pipeline (barrier) and states + acd_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + num_acc_consumer_threads = len(self.epilogue_warp_id) * ( + 2 if use_2cta_instrs else 1 + ) + acd_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_acc_consumer_threads + ) + acd_pipeline = pipeline.PipelineUmmaAsync.create( + barrier_storage=storage.acc_full_mbar_ptr.data_ptr(), + num_stages=self.num_acc_stage, + producer_group=acd_pipeline_producer_group, + consumer_group=acd_pipeline_consumer_group, + cta_layout_vmnk=cluster_layout_vmnk, + defer_sync=True, + ) + + # Pipeline Init: residual TMA load. The dedicated residual warp issues + # the TMA load (single-thread producer arrive) into sC; all epilogue + # warps consume it (one release arrive per warp). No multicast: each CTA + # loads its own output-tile residual. + if cutlass.const_expr(mC_mnl is not None): + c_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + c_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, len(self.epilogue_warp_id) + ) + c_pipeline = pipeline.PipelineTmaAsync.create( + barrier_storage=storage.c_full_mbar_ptr.data_ptr(), + num_stages=self.num_d_stage, + producer_group=c_producer_group, + consumer_group=c_consumer_group, + tx_count=self.num_tma_load_c_bytes, + defer_sync=True, + ) + else: + c_pipeline = None + + # Pipeline Init: Initialize clc_pipeline (CLC fetch async) + # Consumers of CLC response: TMA(1) + cp.async SFA(4) + MMA(1) + Epilogue(4) + # [+ residual(1) when beta != 0] per CTA, * cluster_size; plus sched(1) on + # the first CTA only. The tile-coord query is broadcast to all CTAs in the + # cluster. The residual term must gate on the SAME const as the residual + # warp's CLC-consume block (mC_mnl is not None) or CLC producer/consumer + # counts diverge and the fetch pipeline hangs. + clc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + cluster_size = cute.size(cluster_shape) + num_residual_clc_warps = 1 if cutlass.const_expr(mC_mnl is not None) else 0 + num_clc_consumer_threads = 32 * ( + 1 + + cluster_size + * ( + 1 + + len(self.cpasync_sfa_warp_id) + + len(self.epilogue_warp_id) + + 1 + + num_residual_clc_warps + ) + ) + clc_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_clc_consumer_threads + ) + clc_pipeline = pipeline.PipelineClcFetchAsync.create( + barrier_storage=storage.clc_mbar_ptr.data_ptr(), + num_stages=self.num_clc_stage, + producer_group=clc_pipeline_producer_group, + consumer_group=clc_pipeline_consumer_group, + tx_count=self.num_clc_response_bytes, + cta_layout_vmnk=cluster_layout_vmnk, + defer_sync=True, + ) + + # Pipeline Init: Tensor memory alloc/dealloc barrier init + tmem_alloc_barrier = pipeline.NamedBarrier( + barrier_id=self.tmem_alloc_sync_bar_id, + num_threads=32 * len((self.mma_warp_id, *self.epilogue_warp_id)), + ) + tmem = TmemAllocator( + storage.tmem_holding_buf, + barrier_for_retrieve=tmem_alloc_barrier, + allocator_warp_id=self.epilogue_warp_id[0], + is_two_cta=use_2cta_instrs, + two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar_ptr, + ) + + # Cluster arrive after barrier init + pipeline_init_arrive(cluster_shape_mn=cluster_shape, is_relaxed=True) + + # CLC response buffer pointer + consumer state + clc_response_ptr = storage.clc_response.data_ptr() + clc_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_clc_stage + ) + + # + # Setup smem tensor A/B/SFA/SFB/D + # + # (EPI_TILE_M, EPI_TILE_N, STAGE) + sD = storage.sD.get_tensor( + d_smem_layout_staged.outer, swizzle=d_smem_layout_staged.inner + ) + # (cta_tile_n,) bias row for the shared-memory-staged broadcast. + if cutlass.const_expr(mBias_mnl is not None): + sBias = storage.sBias.get_tensor( + cute.make_layout(self.cta_tile_shape_mnk[1]) + ) + else: + sBias = None + # (EPI_TILE_M, EPI_TILE_N, STAGE) residual tile, same layout as sD. + if cutlass.const_expr(mC_mnl is not None): + sC = storage.sC.get_tensor( + d_smem_layout_staged.outer, swizzle=d_smem_layout_staged.inner + ) + else: + sC = None + # (MMA, MMA_M, MMA_K, STAGE) + sA = storage.sA.get_tensor( + a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner + ) + # (MMA, MMA_N, MMA_K, STAGE) + sB = storage.sB.get_tensor( + b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner + ) + # (MMA, MMA_M, MMA_K, STAGE) + sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged) + # (MMA, MMA_N, MMA_K, STAGE) + sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged) + + # + # Compute multicast mask for A/B/SFA/SFB buffer full + # + a_full_mcast_mask = None + b_full_mcast_mask = None + sfb_full_mcast_mask = None + + if cutlass.const_expr( + effective_is_a_mcast or effective_is_b_mcast or use_2cta_instrs + ): + a_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2 + ) + b_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1 + ) + sfb_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_sfb_vmnk, block_in_cluster_coord_sfb_vmnk, mcast_mode=1 + ) + + # + # Local_tile partition global tensors + # + # (bM, bK, RestM, RestK, RestL) + gA_mkl = cute.local_tile( + mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None) + ) + + # (bN, bK, RestN, RestK, RestL) + gB_nkl = cute.local_tile( + mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None) + ) + # (bN, bK, RestN, RestK, RestL) + gSFB_nkl = cute.local_tile( + mSFB_nkl, + cute.slice_(self.mma_tiler_sfb, (0, None, None)), + (None, None, None), + ) + # (bM, bN, RestM, RestN, RestL) + gD_mnl = cute.local_tile( + mD_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None) + ) + # (bM, bN, RestM, RestN, RestL) - same MNL tiling as gD; the K-axis (N + # of MMA) carries a stride-0 broadcast over self.sfd_vec_size groups. + if cutlass.const_expr(self.gen_sfd): + gSFD_mnl = cute.local_tile( + mSFD_mnl, + cute.slice_(self.mma_tiler, (None, None, 0)), + (None, None, None), + ) + else: + gSFD_mnl = None + # (bM, bN, RestM, RestN, RestL) - residual shares D's MNL tiling; it is a + # full per-element tensor (no broadcast) loaded via TMA, so it is tiled + # exactly like gD. + if cutlass.const_expr(mC_mnl is not None): + gC_mnl = cute.local_tile( + mC_mnl, + cute.slice_(self.mma_tiler, (None, None, 0)), + (None, None, None), + ) + else: + gC_mnl = None + + k_tile_cnt = cute.size(gA_mkl, mode=[3]) + + # + # Partition global tensor for TiledMMA_A/B/C + # + thr_mma = tiled_mma.get_slice(mma_tile_coord_v) + thr_mma_sfb = tiled_mma_sfb.get_slice(mma_tile_coord_v) + # (MMA, MMA_M, MMA_K, RestM, RestK, RestL) + tCgA = thr_mma.partition_A(gA_mkl) + # (MMA, MMA_N, MMA_K, RestN, RestK, RestL) + tCgB = thr_mma.partition_B(gB_nkl) + # (MMA, MMA_N, MMA_K, RestN, RestK, RestL) + tCgSFB = thr_mma_sfb.partition_B(gSFB_nkl) + # (MMA, MMA_M, MMA_N, RestM, RestN, RestL) + tCgD = thr_mma.partition_C(gD_mnl) + # SFD: same partition shape as gD; broadcast strides preserved through partition. + if cutlass.const_expr(self.gen_sfd): + tCgSFD = thr_mma.partition_C(gSFD_mnl) + else: + tCgSFD = None + # Residual: same partition_C chain as the accumulator/D, so the residual + # TMA load walks the output tile exactly like the D store. + if cutlass.const_expr(mC_mnl is not None): + tCgC = thr_mma.partition_C(gC_mnl) + else: + tCgC = None + + # + # Partition global/shared tensor for TMA load B/SFB + # + # TMA load A partition_S/D + a_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape + ) + # ((atom_v, rest_v), STAGE) + # ((atom_v, rest_v), RestM, RestK, RestL) + tAsA, tAgA = cpasync.tma_partition( + tma_atom_a, + block_in_cluster_coord_vmnk[2], + a_cta_layout, + cute.group_modes(sA, 0, 3), + cute.group_modes(tCgA, 0, 3), + ) + # TMA load B partition_S/D + b_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape + ) + # ((atom_v, rest_v), STAGE) + # ((atom_v, rest_v), RestN, RestK, RestL) + tBsB, tBgB = cpasync.tma_partition( + tma_atom_b, + block_in_cluster_coord_vmnk[1], + b_cta_layout, + cute.group_modes(sB, 0, 3), + cute.group_modes(tCgB, 0, 3), + ) + + # TMA load SFB partition_S/D + sfb_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_sfb_vmnk, (0, None, 0, 0)).shape + ) + # ((atom_v, rest_v), STAGE) + # ((atom_v, rest_v), RestN, RestK, RestL) + tBsSFB, tBgSFB = cute.nvgpu.cpasync.tma_partition( + tma_atom_sfb, + block_in_cluster_coord_sfb_vmnk[1], + sfb_cta_layout, + cute.group_modes(sSFB, 0, 3), + cute.group_modes(tCgSFB, 0, 3), + ) + tBsSFB = cute.filter_zeros(tBsSFB) + tBgSFB = cute.filter_zeros(tBgSFB) + + # + # Partition shared/tensor memory tensor for TiledMMA_A/B/C + # + # (MMA, MMA_M, MMA_K, STAGE) + tCrA = tiled_mma.make_fragment_A(sA) + # (MMA, MMA_N, MMA_K, STAGE) + tCrB = tiled_mma.make_fragment_B(sB) + # (MMA, MMA_M, MMA_N, STAGE) + # For overlapping-accum this carries the strided 2-buffer layout so the + # accumulator's second stage aliases the SFA/SFB columns. + tCtAcc_fake = self._make_acc_fake_tensor(tiled_mma, self.mma_tiler) + + # + # Cluster wait before tensor memory alloc + # + pipeline_init_wait(cluster_shape_mn=cluster_shape) + + # + # Specialized TMA load A/B/SFA/SFB warp + # + if warp_idx == self.tma_warp_id: + # + # Persistent tile scheduling loop + # + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + + ab_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_ab_stage + ) + + while work_tile.is_valid_tile: + # Get tile coord from tile scheduler + cur_tile_coord = work_tile.tile_idx + mma_tile_coord_mnl = ( + cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape), + cur_tile_coord[1], + cur_tile_coord[2], + ) + + # + # Slice to per mma tile index + # + # ((atom_v, rest_v), RestK) + tAgA_slice = tAgA[ + (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2]) + ] + # ((atom_v, rest_v), RestK) + tBgB_slice = tBgB[ + (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2]) + ] + + # Special-case for N=64: SFB scale-factor slicing to match Tensor Memory packing + slice_n = mma_tile_coord_mnl[1] + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64): + slice_n = mma_tile_coord_mnl[1] // 2 + # ((atom_v, rest_v), RestK) + tBgSFB_slice = tBgSFB[(None, slice_n, None, mma_tile_coord_mnl[2])] + + # Peek (try_wait) shared AB/SFB buffer empty + ab_producer_state.reset_count() + peek_ab_empty_status = ab_sfa_pipeline.producer_try_acquire( + ab_producer_state + ) + + # + # Tma load loop with increment_coord for im2col K traversal + # Traversal order: S innermost -> R -> T -> C outermost + # + # A's RestK shape from im2col TMA is (C_tiles, S, R, T) where + # C_tiles is leftmost. To traverse S->R->T->C, we build a + # permuted traversal shape and remap coords when indexing. + k_shape = cute.shape(tAgA_slice, mode=1) # (C_tiles, S, R, T) + k_C_tiles = cute.size(k_shape, mode=0) + k_S = cute.size(k_shape, mode=1) + k_R = cute.size(k_shape, mode=2) + k_T = cute.size(k_shape, mode=3) + # SFB boxes per filter position: one box is a whole SF K tile, so a + # position holds that many fewer boxes than A/B tiles. Rounding up + # matches the span the host padded SFB to. + sfb_C_tiles = cute.ceil_div(k_C_tiles, self.ab_k_tiles_per_sf_k_tile) + # Traversal shape in S->R->T->C order (colexicographic on this) + trav_shape = (k_S, k_R, k_T, k_C_tiles) + trav_coord = cute.repeat_like(0, trav_shape) + + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + # TMA producer path on the merged pipeline (leader does + # arrive_and_expect_tx, peer does nothing). + ab_sfa_pipeline.producer_acquire_tma( + ab_producer_state, peek_ab_empty_status + ) + + # Remap (s,r,t,c) -> (c,s,r,t) to match A/B RestK layout + ab_coord = ( + trav_coord[3], # C_tiles + trav_coord[0], # S + trav_coord[1], # R + trav_coord[2], # T + ) + + # TMA load A (use multi-dim coord for im2col K indexing) + cute.copy( + tma_atom_a, + tAgA_slice[(None, ab_coord)], + tAsA[(None, ab_producer_state.index)], + tma_bar_ptr=ab_sfa_pipeline.producer_get_barrier( + ab_producer_state + ), + mcast_mask=a_full_mcast_mask, + ) + # TMA load B (multi-dim K layout, use same coord) + cute.copy( + tma_atom_b, + tBgB_slice[(None, ab_coord)], + tBsB[(None, ab_producer_state.index)], + tma_bar_ptr=ab_sfa_pipeline.producer_get_barrier( + ab_producer_state + ), + mcast_mask=b_full_mcast_mask, + ) + # TMA load SFB: compute flat K tile index from (c,s,r,t) + # SFB uses flat layout with C->S->R->T physical order. The + # channel coordinate is on the SF cadence, so the A/B tiles + # sharing one chunk name the same box and each stages it whole. + sfb_flat_k = ( + ab_coord[0] // self.ab_k_tiles_per_sf_k_tile + + ab_coord[1] * sfb_C_tiles + + ab_coord[2] * k_S * sfb_C_tiles + + ab_coord[3] * k_R * k_S * sfb_C_tiles + ) + cute.copy( + tma_atom_sfb, + tBgSFB_slice[(None, sfb_flat_k)], + tBsSFB[(None, ab_producer_state.index)], + tma_bar_ptr=ab_sfa_pipeline.producer_get_barrier( + ab_producer_state + ), + mcast_mask=sfb_full_mcast_mask, + ) + + # Peek (try_wait) shared AB/SFB buffer empty for next iteration + ab_producer_state.advance() + peek_ab_empty_status = cutlass.Boolean(1) + if ab_producer_state.count < k_tile_cnt: + peek_ab_empty_status = ab_sfa_pipeline.producer_try_acquire( + ab_producer_state + ) + + # Advance traversal coord in S->R->T->C order + trav_coord = cute.increment_coord(trav_coord, trav_shape) + + # + # Advance to next tile (CLC consumer) + # + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + + # + # Wait shared AB/SFB buffer empty + # + ab_sfa_pipeline.producer_tail(ab_producer_state) + + # + # Specialized cp.async SFA warp + # + if ( + warp_idx >= self.cpasync_sfa_warp_id[0] + and warp_idx <= self.cpasync_sfa_warp_id[-1] + ): + # + # Setup SFA CPASYNC copy atom + # + sfa_atom_copy = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(), + mSFA_mkl.element_type, + num_bits_per_copy=32, + ) + tidx_in_warpgroup = tidx % 128 + + # SFA predicate: dynamically computed per cp.async for conv3x3 + sfa_predicate_tensor = cute.make_rmem_tensor( + cute.make_layout((1,)), + cutlass.Boolean, + ) + + # Identity M offset (no permutation) + CTA offset within cluster + cta_m_offset = mma_tile_coord_v * self.cta_tile_shape_mnk[0] + sfa_m_offset = ( + cta_m_offset + + 8 * (tidx_in_warpgroup // 32) + + 32 * ((tidx_in_warpgroup % 32) // 8) + + (tidx_in_warpgroup % 8) + ) + + # Slice sSFA for this thread's M-row and sub-row + tAsSFA = sSFA[ + ( + ( + ( + ( + 8 * (tidx_in_warpgroup // 32) + (tidx_in_warpgroup % 8), + (tidx_in_warpgroup % 32) // 8, + ), + None, + ), + None, + ), + None, + None, + None, + ) + ] + + # SFA global tensor: shape (MN, SF_K, L) with K contiguous + # For conv3x3 with arbitrary C, each cp.async may load from a different + # input element (different M address) depending on the filter position. + # When C < K_gemm_tile, a single K tile spans multiple filter positions. + + # Precompute conv spatial constants + K_gemm_tile = self.mma_tiler[2][0] + ZPQ = output_Z * output_P * output_Q + PQ = output_P * output_Q + DHW = input_D * input_H * input_W + HW = input_H * input_W + CTA_M = self.cta_tile_shape_mnk[0] + # SF blocks one SF K tile holds. The SF tile widens to a whole atom + # when the A/B K tile is narrower, and every A/B tile sharing it stages + # the whole thing. + sf_k_per_ktile = self.sf_tile_k // self.sf_vec_size + # C tiles per filter position, rounded up: a partial trailing tile still + # consumes a k_tile, and its tail reads zero-filled A and B from the TMA. + c_tiles_per_fpos = (input_C + K_gemm_tile - 1) // K_gemm_tile + # SF blocks a pixel's row holds. The host rounds it up to a whole 4-block + # atom and zero-fills the tail, so a group on the last one is in bounds. + sfa_sf_k_extent = mSFA_mkl.shape[1] + sfa_last_group = sfa_sf_k_extent - 4 + # Per-carry address deltas for the linear im2col walk: the unclamped + # gmem M address decomposes as m_base + t*dil_d*HW + r*dil_h*W + + # s*dil_w, so each carry level advances one running delta by a + # precomputed step. Steps carry the SFA row stride, so the loop + # adds element offsets directly. + sfa_step_s = dil_w * sfa_sf_k_extent + sfa_step_r = (dil_h * input_W - conv_S * dil_w) * sfa_sf_k_extent + sfa_step_t = (dil_d * HW - conv_R * dil_h * input_W) * sfa_sf_k_extent + + # + # Persistent tile scheduling loop + # + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + + sfa_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_ab_stage + ) + + while work_tile.is_valid_tile: + # Get tile coord from tile scheduler + cur_tile_coord = work_tile.tile_idx + mma_tile_coord_mnl = ( + cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape), + cur_tile_coord[1], + cur_tile_coord[2], + ) + + # Compute this thread's global M index (output pixel) + # Use mma_tile_coord_mnl[0] (M tile index without V) * full MMA tile M + # because sfa_m_offset already includes the V (CTA-within-pair) offset + cta_m_tile_offset = ( + mma_tile_coord_mnl[0] * CTA_M * cute.size(tiled_mma.thr_id.shape) + ) + m_global = cta_m_tile_offset + sfa_m_offset + + # Decompose m_global -> (n, z, p, q) output coordinates + n_idx = m_global // ZPQ + zpq_rem = m_global % ZPQ + z_idx = zpq_rem // PQ + pq_rem = zpq_rem % PQ + p_idx = pq_rem // output_Q + q_idx = pq_rem % output_Q + + # Check if m_global is within valid output range + m_valid = m_global < (input_N * ZPQ) + + # Initialize base spatial coordinates for filter pos (t=0, r=0, s=0) + d_in_base = z_idx * stride_d - pad_d + h_in_base = p_idx * stride_h - pad_h + w_in_base = q_idx * stride_w - pad_w + n_clamped = n_idx if m_valid else 0 + + # Maintain (s,r,t,c) coords for S->R->T->C traversal order + # matching TMA warp's coord_iter order + sfa_s_idx = 0 + sfa_r_idx = 0 + sfa_t_idx = 0 + sfa_c_tile_idx = 0 + sfa_m_base = ( + n_clamped * DHW + d_in_base * HW + h_in_base * input_W + w_in_base + ) * sfa_sf_k_extent + sfa_delta = cutlass.Int32(0) + # Running unclamped coords for the predicate, advanced with the + # same carry chain as the delta. + sfa_d_in = d_in_base + 0 + sfa_h_in = h_in_base + 0 + sfa_w_in = w_in_base + 0 + + # Peek (try_wait) SFA buffer empty + sfa_producer_state.reset_count() + peek_sfa_empty_status = cutlass.Boolean(1) + if sfa_producer_state.count < k_tile_cnt: + peek_sfa_empty_status = ab_sfa_pipeline.producer_try_acquire( + sfa_producer_state + ) + + # + # Shift-by-K cross-CTA commit ringbuffer (K=2). Per iter, + # commit the prior copies as a cp.async group and wait_group(K); + # then remote-arrive on leader's sfa_full for the stage from + # K iters ago. Guarantees that cp.async has landed globally + # before leader MMA sees the barrier full. + # + sfa_shift_K = 2 + pending_a = cutlass.Int32(0) + pending_b = cutlass.Int32(0) + + # + # CPASYNC SFA load loop with coordinate-based im2col + # Traversal order: S->R->T->C (matching TMA warp) + # + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + # Conditionally wait for SFA buffer empty + ab_sfa_pipeline.producer_acquire( + sfa_producer_state, peek_sfa_empty_status + ) + + cur_stage = sfa_producer_state.index + + tAsSFA_ktile = tAsSFA[(None, None, None, None, cur_stage)] + + # Predicate: m_valid AND spatial bounds check + sfa_pred_val = cutlass.Boolean(0) + if m_valid: + if sfa_d_in >= 0 and sfa_d_in < input_D: + if sfa_h_in >= 0 and sfa_h_in < input_H: + if sfa_w_in >= 0 and sfa_w_in < input_W: + sfa_pred_val = cutlass.Boolean(1) + + # A predicated-off row reads address 0, always in bounds; + # its predicate suppresses the actual copy. + sfa_row_off = (sfa_m_base + sfa_delta) if sfa_pred_val else 0 + + # SF K base offset for current C tile within this fpos. A/B + # tiles sharing one SF tile take the same base: they sit in one + # filter position, so one transfer covers them all. + sf_k_base = ( + sfa_c_tile_idx // self.ab_k_tiles_per_sf_k_tile + ) * sf_k_per_ktile + + # 4-byte cp.async SFA -- all SI share same fpos. + # Each thread must cover every SI in [0, num_sfa_cpasync) exactly + # once so all SF-K groups of its row get written. XOR (q ^ i) is a + # complete cover only when num_sfa_cpasync is a power of 2; for a + # non-power-of-2 count (e.g. 3 when cta_tile_k == 3 * mma_inst_shape_k) + # it drops SI values on some quarter-warps, leaving SF-K groups + # unwritten. Additive rotation (q + i) % num is a complete cyclic + # cover for any count and keeps quarter-warps on distinct groups + # per timestep. + for i in range(self.num_sfa_cpasync): + SI = ((tidx_in_warpgroup % 32) // 8 + i) % self.num_sfa_cpasync + + # One 4-byte cp.async moves 4 contiguous SF blocks (4 x 1-byte + # E4M3 = the 4-byte transfer). SI indexes which group of 4, + # so its SF-block base is SI * 4 and the copy shape is (4,). + sf_k_smem = SI * 4 + local_sf_k = sf_k_base + (sf_k_smem % sf_k_per_ktile) + # A partial trailing tile reaches past the scale factors this + # pixel owns. Clamping the group start to the row's last whole + # atom keeps it in bounds and finite, which the MMA needs since + # it applies the scale before seeing the zero operand. + local_sf_k = ( + local_sf_k + if local_sf_k <= sfa_last_group + else sfa_last_group + ) + + sfa_predicate_tensor[0] = sfa_pred_val + + # Gmem: this fpos's row offset plus the local SF-K offset + tAgSFA_slice_ptr = mSFA_mkl.iterator + ( + sfa_row_off + local_sf_k + ) + tAgSFA_slice = cute.make_tensor( + tAgSFA_slice_ptr, layout=cute.make_layout((4,)) + ) + + # Smem: write to swizzled position. Adjacent SF-block + # groups are 512 bytes apart in the SFA smem layout: + # CTA_M (128 rows) x 4 bytes per group = 512. + tAsSFA_slice_ptr = tAsSFA_ktile.iterator + 512 * SI + tAsSFA_slice = cute.make_tensor( + tAsSFA_slice_ptr, cute.make_layout((4,)) + ) + + cute.copy_atom_call( + sfa_atom_copy, + tAgSFA_slice, + tAsSFA_slice, + pred=sfa_predicate_tensor, + ) + + # Increment in S->R->T->C order (S innermost, C outermost) + sfa_s_idx = sfa_s_idx + 1 + sfa_delta = sfa_delta + sfa_step_s + sfa_w_in = sfa_w_in + dil_w + if sfa_s_idx >= conv_S: + sfa_s_idx = 0 + sfa_w_in = w_in_base + 0 + sfa_r_idx = sfa_r_idx + 1 + sfa_delta = sfa_delta + sfa_step_r + sfa_h_in = sfa_h_in + dil_h + if sfa_r_idx >= conv_R: + sfa_r_idx = 0 + sfa_h_in = h_in_base + 0 + sfa_t_idx = sfa_t_idx + 1 + sfa_delta = sfa_delta + sfa_step_t + sfa_d_in = sfa_d_in + dil_d + if sfa_t_idx >= conv_T: + sfa_t_idx = 0 + sfa_d_in = d_in_base + 0 + sfa_delta = 0 + sfa_c_tile_idx = sfa_c_tile_idx + 1 + if sfa_c_tile_idx >= c_tiles_per_fpos: + sfa_c_tile_idx = 0 + + # Commit this iter's copies as one cp.async group and + # gate on <=K inflight. After wait_group(K), the cp.async + # for the stage from K iters ago is guaranteed complete + # on every CTA in the cluster. + ab_sfa_pipeline.producer_cp_async_commit() + ab_sfa_pipeline.producer_cp_async_wait(sfa_shift_K) + + # Shift-by-K: arrive on stage from K iters ago. + if k_tile >= sfa_shift_K: + ab_sfa_pipeline.producer_arrive_remote(pending_a) + + # Advance ringbuffer (oldest <- 2nd oldest <- current). + pending_a = pending_b + pending_b = cur_stage + + # Peek (try_wait) SFA buffer empty for next iteration + sfa_producer_state.advance() + peek_sfa_empty_status = cutlass.Boolean(1) + if sfa_producer_state.count < k_tile_cnt: + peek_sfa_empty_status = ab_sfa_pipeline.producer_try_acquire( + sfa_producer_state + ) + + # + # Shift-by-K tail: drain all inflight cp.async, then arrive on + # the K stages whose arrives were deferred. + # + ab_sfa_pipeline.producer_cp_async_wait(0) + if k_tile_cnt >= sfa_shift_K: + ab_sfa_pipeline.producer_arrive_remote(pending_a) + if k_tile_cnt >= 1: + ab_sfa_pipeline.producer_arrive_remote(pending_b) + + # + # Advance to next tile (CLC consumer) + # + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + + # + # Wait SFA buffer empty + # + ab_sfa_pipeline.producer_tail(sfa_producer_state) + + # + # Specialized residual G2S warp (only spawned when beta != 0). + # + # The residual load owns a dedicated warp, separate from the TMA warp that + # issues A/B/SFB. Producer and epilogue consumer still sit on physically + # separate warps (required: a same-warp producer/consumer pair on the + # residual barrier self-deadlocks the epilogue). It runs its own persistent + # tile-scheduler loop and consumes one CLC response per tile, so it is + # counted in num_clc_consumer_threads under the SAME const gate + # (mC_mnl is not None). + # + if warp_idx == self.residual_warp_id: + if cutlass.const_expr(mC_mnl is not None): + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + + # The gmem/smem TMA partition is thread-invariant, so recompute + # the exact handles the epilogue consumer expects. + ( + tma_atom_c, + bGS_sC, + bGS_gC_partitioned, + ) = self.epilog_gmem_copy_and_partition( + tidx, tma_atom_c, tCgC, epi_tile, sC + ) + c_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_d_stage + ) + # Overlapping-accum walks output subtiles back-to-front on the + # phase-0 tile; mirror the epilogue's acc-consumer phase with a + # shadow state so the producer stages subtiles in the same order + # the consumer reads them. + if cutlass.const_expr(self.overlapping_accum): + c_acc_shadow_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_acc_stage + ) + + while work_tile.is_valid_tile: + cur_tile_coord = work_tile.tile_idx + mma_tile_coord_mnl = ( + cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape), + cur_tile_coord[1], + cur_tile_coord[2], + ) + + # + # Issue this output tile's residual TMA loads (gmem -> smem), + # one per output N-subtile, walking subtiles in lockstep with + # the epilogue consumer. Whole-warp producer arrive (no + # elect_one): PipelineTmaAsync's single-thread producer group + # self-elects the signaling lane. + # + bGS_gC = bGS_gC_partitioned[ + ( + None, + None, + None, + *mma_tile_coord_mnl, + ) + ] + bGS_gC = cute.group_modes(bGS_gC, 1, cute.rank(bGS_gC)) + subtile_cnt = cute.size(bGS_gC.shape, mode=[1]) + if cutlass.const_expr(self.overlapping_accum): + reverse_subtile = c_acc_shadow_state.phase == 0 + for subtile_idx in cutlass.range(subtile_cnt): + real_subtile_idx = subtile_idx + if cutlass.const_expr(self.overlapping_accum): + if reverse_subtile: + real_subtile_idx = ( + self.cta_tile_shape_mnk[1] // self.epi_tile_n + - 1 + - subtile_idx + ) + # The ring slot comes from the pipeline state, the same + # place the barrier below and the epilogue consumer take + # theirs, so the ring holds at any stage count. + c_pipeline.producer_acquire(c_producer_state) + cute.copy( + tma_atom_c, + bGS_gC[(None, real_subtile_idx)], + bGS_sC[(None, c_producer_state.index)], + tma_bar_ptr=c_pipeline.producer_get_barrier( + c_producer_state + ), + ) + c_producer_state.advance() + if cutlass.const_expr(self.overlapping_accum): + c_acc_shadow_state.advance() + + # + # Advance to next tile (CLC consumer) + # + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + + # + # Drain the residual TMA load pipeline before exiting. + # + c_pipeline.producer_tail(c_producer_state) + + # + # Specialized scheduler warp (drives CLC fetch, only first CTA in cluster) + # + if warp_idx == self.sched_warp_id and is_first_cta_in_cluster: + clc_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.ProducerConsumer, self.num_clc_stage + ) + + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + while work_tile.is_valid_tile: + clc_pipeline.producer_acquire(clc_producer_state) + mbarrier_addr = clc_pipeline.producer_get_barrier(clc_producer_state) + tile_sched.advance_to_next_work(mbarrier_addr) + clc_producer_state.advance() + + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + clc_pipeline.producer_tail(clc_producer_state) + + # + # Specialized MMA warp + # + if warp_idx == self.mma_warp_id: + # + # Bar sync for retrieve tensor memory ptr from shared mem + # + tmem.wait_for_alloc() + + # + # Retrieving tensor memory ptr and make accumulator/SFA/SFB tensor + # + acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + # Make accumulator tmem tensor + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout) + + # Make SFA tmem tensor + sfa_tmem_ptr = cute.recast_ptr( + acc_tmem_ptr + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base), + dtype=self.sf_dtype, + ) + # (MMA, MMA_M, MMA_K) + tCtSFA_layout = blockscaled_utils.make_tmem_layout_sfa( + tiled_mma, + self.flat_mma_tiler_sf, + self.sf_vec_size, + cute.slice_(sfa_smem_layout_staged, (None, None, None, 0)), + ) + tCtSFA = cute.make_tensor(sfa_tmem_ptr, tCtSFA_layout) + + # Make SFB tmem tensor + sfb_tmem_ptr = cute.recast_ptr( + acc_tmem_ptr + + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base) + + tcgen05.find_tmem_tensor_col_offset(tCtSFA), + dtype=self.sf_dtype, + ) + # (MMA, MMA_N, MMA_K) + tCtSFB_layout = blockscaled_utils.make_tmem_layout_sfb( + tiled_mma, + self.flat_mma_tiler_sf, + self.sf_vec_size, + cute.slice_(sfb_smem_layout_staged, (None, None, None, 0)), + ) + tCtSFB = cute.make_tensor(sfb_tmem_ptr, tCtSFB_layout) + # + # Partition for S2T copy of SFA/SFB + # + ( + tiled_copy_s2t_sfa, + tCsSFA_compact_s2t, + tCtSFA_compact_s2t, + ) = self.mainloop_s2t_copy_and_partition(sSFA, tCtSFA) + ( + tiled_copy_s2t_sfb, + tCsSFB_compact_s2t, + tCtSFB_compact_s2t, + ) = self.mainloop_s2t_copy_and_partition(sSFB, tCtSFB) + + # + # Persistent tile scheduling loop + # + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + + ab_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_ab_stage + ) + sfa_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_ab_stage + ) + acc_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_acc_stage + ) + + # Filter positions one channel tile sweeps. The K loop walks S->R->T->C + # with C outermost, so the chunk part below has to come from the channel + # tile: the running k_tile count keeps advancing across positions. + if cutlass.const_expr(self.ab_k_tiles_per_sf_k_tile > 1): + sf_fpos_per_c_tile = conv_T * conv_R * conv_S + + while work_tile.is_valid_tile: + # Get tile coord from tile scheduler + cur_tile_coord = work_tile.tile_idx + mma_tile_coord_mnl = ( + cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape), + cur_tile_coord[1], + cur_tile_coord[2], + ) + + # Channel tile and its filter-position counter, reset per work tile. + if cutlass.const_expr(self.ab_k_tiles_per_sf_k_tile > 1): + sf_c_tile_idx = cutlass.Int32(0) + sf_fpos_idx = cutlass.Int32(0) + + # Set tensor memory buffer for current tile + # (MMA, MMA_M, MMA_N) + # Overlapping-accum squeezes two acc buffers into 512-col TMEM by + # aliasing the second buffer onto the SFA/SFB columns. The producer + # writes the buffer the consumer is NOT currently draining, which is + # acc_producer_state.phase ^ 1 (producer phase inits to 1, consumer to + # 0, so tile 0 picks stage 0 to match the consumer). + if cutlass.const_expr(self.overlapping_accum): + acc_stage_index = acc_producer_state.phase ^ 1 + else: + acc_stage_index = acc_producer_state.index + tCtAcc = tCtAcc_base[(None, None, None, acc_stage_index)] + + # Peek (try_wait) shared AB/SFB + SFA buffer full for k_tile = 0 + ab_consumer_state.reset_count() + sfa_consumer_state.reset_count() + peek_ab_full_status = cutlass.Boolean(1) + if ab_consumer_state.count < k_tile_cnt and is_leader_cta: + peek_ab_full_status = ab_sfa_pipeline.consumer_try_wait( + ab_consumer_state + ) + # Merged pipeline: sfa shares ab mbarrier, no separate try_wait needed. + + # + # Wait for accumulator buffer empty + # + if is_leader_cta: + acd_pipeline.producer_acquire(acc_producer_state) + + # Special-case for N=192/N=64: shift the SFB Tensor Memory pointer to match its packing + tCtSFB_mma = tCtSFB + if cutlass.const_expr(self.cta_tile_shape_mnk[1] == 192): + # If this is an ODD tile, shift the Tensor Memory start address for cta_tile_shape_n=192 case by two words (ignores first 64 columns of SFB) + offset = ( + cutlass.Int32(2) + if mma_tile_coord_mnl[1] % 2 == 1 + else cutlass.Int32(0) + ) + shifted_ptr = cute.recast_ptr( + acc_tmem_ptr + + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base) + + tcgen05.find_tmem_tensor_col_offset(tCtSFA) + + offset, + dtype=self.sf_dtype, + ) + tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout) + elif cutlass.const_expr(self.cta_tile_shape_mnk[1] == 64): + # Move in increments of 64 columns of SFB + offset = cutlass.Int32((mma_tile_coord_mnl[1] % 2) * 2) + shifted_ptr = cute.recast_ptr( + acc_tmem_ptr + + tcgen05.find_tmem_tensor_col_offset(tCtAcc_base) + + tcgen05.find_tmem_tensor_col_offset(tCtSFA) + + offset, + dtype=self.sf_dtype, + ) + tCtSFB_mma = cute.make_tensor(shifted_ptr, tCtSFB_layout) + + # + # Reset the ACCUMULATE field for each tile + # + tiled_mma.set(tcgen05.Field.ACCUMULATE, False) + + # + # Mma mainloop + # + for k_tile in cutlass.range(k_tile_cnt, unroll=1): + if cutlass.const_expr(self.ab_k_tiles_per_sf_k_tile > 1): + sf_kblock_base = ( + sf_c_tile_idx % self.ab_k_tiles_per_sf_k_tile + ) * self.mma_inst_tile_k + else: + sf_kblock_base = 0 + + if is_leader_cta: + # Merged pipeline: single consumer_wait covers both TMA and cp.async producers. + ab_sfa_pipeline.consumer_wait( + ab_consumer_state, peek_ab_full_status + ) + + # Copy SFA/SFB from smem to tmem + sfa_s2t_stage_coord = ( + None, + None, + None, + None, + sfa_consumer_state.index, + ) + sfb_s2t_stage_coord = ( + None, + None, + None, + None, + ab_consumer_state.index, + ) + tCsSFA_compact_s2t_staged = tCsSFA_compact_s2t[ + sfa_s2t_stage_coord + ] + tCsSFB_compact_s2t_staged = tCsSFB_compact_s2t[ + sfb_s2t_stage_coord + ] + cute.copy( + tiled_copy_s2t_sfa, + tCsSFA_compact_s2t_staged, + tCtSFA_compact_s2t, + ) + cute.copy( + tiled_copy_s2t_sfb, + tCsSFB_compact_s2t_staged, + tCtSFB_compact_s2t, + ) + + # Block-scaled MMA: acc += (A * SFA) * (B * SFB) + num_kblocks = cute.size(tCrA, mode=[2]) + for kblock_idx in cutlass.range(num_kblocks, unroll_full=True): + kblock_coord = ( + None, + None, + kblock_idx, + ab_consumer_state.index, + ) + + # Set SFA/SFB tensor to tiled_mma. The K blocks this + # tile owns start at its own part of the shared chunk + # (offset 0 for NVFP4, whose ratio is 1). + sf_kblock_coord = ( + None, + None, + sf_kblock_base + kblock_idx, + ) + tiled_mma.set( + tcgen05.Field.SFA, + tCtSFA[sf_kblock_coord].iterator, + ) + tiled_mma.set( + tcgen05.Field.SFB, + tCtSFB_mma[sf_kblock_coord].iterator, + ) + + cute.gemm( + tiled_mma, + tCtAcc, + tCrA[kblock_coord], + tCrB[kblock_coord], + tCtAcc, + ) + + # Enable accumulate on tCtAcc after first kblock + tiled_mma.set(tcgen05.Field.ACCUMULATE, True) + + # Merged pipeline: single consumer_release covers both TMA and cp.async producers. + ab_sfa_pipeline.consumer_release(ab_consumer_state) + + # Advance the channel tile once the whole filter has been + # swept, mirroring the S->R->T->C walk of the load warps. + if cutlass.const_expr(self.ab_k_tiles_per_sf_k_tile > 1): + sf_fpos_idx = sf_fpos_idx + 1 + if sf_fpos_idx >= sf_fpos_per_c_tile: + sf_fpos_idx = cutlass.Int32(0) + sf_c_tile_idx = sf_c_tile_idx + 1 + + # Peek (try_wait) shared AB/SFB + SFA buffer full for k_tile+1 + ab_consumer_state.advance() + sfa_consumer_state.advance() + + peek_ab_full_status = cutlass.Boolean(1) + if ab_consumer_state.count < k_tile_cnt: + if is_leader_cta: + peek_ab_full_status = ab_sfa_pipeline.consumer_try_wait( + ab_consumer_state + ) + # Merged pipeline: sfa shares ab mbarrier, try_wait above handles both. + + # + # Async arrive accumulator buffer full + # + if is_leader_cta: + acd_pipeline.producer_commit(acc_producer_state) + acc_producer_state.advance() + + # + # Advance to next tile (CLC consumer) + # + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + + # + # Wait for accumulator buffer empty + # + acd_pipeline.producer_tail(acc_producer_state) + + # + # Specialized epilogue warps + # + if warp_idx <= self.epilogue_warp_id[-1]: + # + # Alloc tensor memory buffer + # + tmem.allocate(self.num_tmem_alloc_cols) + + # + # Bar sync for retrieve tensor memory ptr from shared memory + # + tmem.wait_for_alloc() + + # + # Retrieving tensor memory ptr and make accumulator tensor + # + acc_tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_base = cute.make_tensor(acc_tmem_ptr, tCtAcc_fake.layout) + + # + # Partition for epilogue + # + epi_tidx = tidx + ( + tiled_copy_t2r, + tTR_tAcc_base, + tTR_rAcc, + ) = self.epilog_tmem_copy_and_partition( + epi_tidx, tCtAcc_base, tCgD, epi_tile, use_2cta_instrs + ) + + tTR_rD = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + tiled_copy_r2s, tRS_rD, tRS_sD = self._epilogue_smem_copy_and_partition( + tiled_copy_t2r, tTR_rD, epi_tidx, sD + ) + ( + tma_atom_d, + bSG_sD, + bSG_gD_partitioned, + ) = self.epilog_gmem_copy_and_partition( + epi_tidx, tma_atom_d, tCgD, epi_tile, sD + ) + + # Residual smem->reg partition: mirrors the accumulator's T2R + # fragment so each thread's residual register tile lines up with its + # acc tile. The gmem->smem TMA producer runs on the TMA warp; the + # epilogue only consumes the staged smem tile here. + if cutlass.const_expr(mC_mnl is not None): + tTR_rC = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + ( + tiled_copy_s2r_c, + tSR_rC, + tSR_sC, + ) = self.epilog_smem_load_copy_and_partition( + tiled_copy_t2r, tTR_rC, epi_tidx, sC + ) + else: + tTR_rC = None + tiled_copy_s2r_c = None + tSR_rC = None + tSR_sC = None + + # SFD partition: same epi-tile/T2R structure as gD, but written + # directly to GMEM via STG (no SMEM/TMA path). + if cutlass.const_expr(self.gen_sfd): + thr_copy_t2r = tiled_copy_t2r.get_slice(epi_tidx) + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL) + gSFD_epi = cute.flat_divide( + tCgSFD[((None, None), 0, 0, None, None, None)], epi_tile + ) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL) + tTR_gSFD_partitioned = thr_copy_t2r.partition_D(gSFD_epi) + # The SFD store is a plain STG whose partition rounds M up to a + # tile multiple, so the last m-tile has overhang rows with no + # backing storage. Build an identity-coordinate mirror through + # the D path (which has no stride-0 broadcast mode), partitioned + # identically, so the store site can recover each thread's M + # coordinate and predicate the STG against the real M extent. + sfd_ceil_m, sfd_ceil_n, _ = cute.ceil_div( + mD_mnl.shape, (self.mma_tiler[0], self.mma_tiler[1], 1) + ) + mcSFD = cute.make_identity_tensor( + ( + cute.size(sfd_ceil_m) * self.mma_tiler[0], + cute.size(sfd_ceil_n) * self.mma_tiler[1], + 1, + ) + ) + gcSFD = cute.local_tile( + mcSFD, + cute.slice_(self.mma_tiler, (None, None, 0)), + (None, None, None), + ) + tCgcSFD = thr_mma.partition_C(gcSFD) + gcSFD_epi = cute.flat_divide( + tCgcSFD[((None, None), 0, 0, None, None, None)], epi_tile + ) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL) + tTR_cSFD_partitioned = thr_copy_t2r.partition_D(gcSFD_epi) + else: + tTR_gSFD_partitioned = None + tTR_cSFD_partitioned = None + + # Bias read view: the staged row carries the same epi-tile/T2R + # structure as the accumulator, so each epilogue thread's bias + # register fragment lines up with its acc fragment. It is tile + # invariant, since every tile restages the row in place. + if cutlass.const_expr(mBias_mnl is not None): + # (T2R, T2R_M, T2R_N, SUBTILE) + tTR_sBias = self._epilogue_bias_smem_partition( + tiled_copy_t2r, epi_tidx, sBias, epi_tile + ) + # smem-load atom: each thread reads its own N columns back from the + # CTA-staged sBias row. + simt_atom_bias = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), mBias_mnl.element_type + ) + # cp.async transfers at least 32 bits, so each lane moves a + # 32-bit-wide vector of bias elements (2 bf16/fp16, 1 fp32). + bias_elems_per_copy = 32 // mBias_mnl.element_type.width + bias_g2s_atom = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(), + mBias_mnl.element_type, + num_bits_per_copy=32, + ) + else: + tTR_sBias = None + simt_atom_bias = None + bias_g2s_atom = None + bias_elems_per_copy = None + + # + # Persistent tile scheduling loop + # + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + + acc_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_acc_stage + ) + + # Pipeline Init: Threads/warps participating in tma store pipeline + d_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + 32 * len(self.epilogue_warp_id), + ) + d_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.num_d_stage, + producer_group=d_producer_group, + ) + + # The residual producer state lives on the dedicated residual warp; + # the epilogue only tracks the consumer side of the residual load + # pipeline. + if cutlass.const_expr(mC_mnl is not None): + c_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_d_stage + ) + else: + c_consumer_state = None + + while work_tile.is_valid_tile: + # Get tile coord from tile scheduler + cur_tile_coord = work_tile.tile_idx + mma_tile_coord_mnl = ( + cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape), + cur_tile_coord[1], + cur_tile_coord[2], + ) + # + # Slice to per mma tile index + # + # ((ATOM_V, REST_V), EPI_M, EPI_N) + bSG_gD = bSG_gD_partitioned[ + ( + None, + None, + None, + *mma_tile_coord_mnl, + ) + ] + if cutlass.const_expr(self.gen_sfd): + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N) + tTR_gSFD = tTR_gSFD_partitioned[ + ( + None, + None, + None, + None, + None, + *mma_tile_coord_mnl, + ) + ] + # Identity-coordinate mirror, sliced identically, for the + # STG M-bound predicate at the SFD store site. + tTR_cSFD = tTR_cSFD_partitioned[ + ( + None, + None, + None, + None, + None, + *mma_tile_coord_mnl, + ) + ] + if cutlass.const_expr(mBias_mnl is not None): + # The CTA cp.async's this tile's contiguous cta_tile_n bias + # values (K stride 1) into sBias, then syncs so every + # epilogue thread can load its columns back from smem. One value per + # output channel is fetched once and broadcast to all M rows. + cta_n = self.cta_tile_shape_mnk[1] + n_base = mma_tile_coord_mnl[1] * cta_n + # Each lane cp.async's a contiguous 32-bit vector + # (bias_elems_per_copy elements) of this tile's bias row. + # Layout (elems, lanes) with column-major stride so lane t + # owns the contiguous block [t*elems : (t+1)*elems); the K + # (output-channel) axis of mBias_mnl has stride 1, so the + # segment is contiguous. n_active lanes cover cta_n. + n_active = cta_n // bias_elems_per_copy + row_layout = cute.make_layout( + (bias_elems_per_copy, n_active), + stride=(1, bias_elems_per_copy), + ) + # cp.async needs 32-bit source/dest alignment; the tile + # base (n_base is a multiple of cta_tile_n) keeps both on a + # 4-byte boundary, so re-annotate the pointers to satisfy + # the 32-bit atom. + gBias_row = cute.make_tensor( + (mBias_mnl.iterator + n_base).align(min_align=4), + row_layout, + ) + sBias_tiled = cute.make_tensor( + sBias.iterator.align(min_align=4), row_layout + ) + # Predicated cp.async: a CTA N-tile rounds up to cta_tile_n, + # but K (= output channels = GEMM-N) need not divide it, so + # the tail lanes address bias columns n >= K that have + # no backing storage. Guard each lane's contiguous vector on + # its base channel: in-bounds lanes cp.async from gmem, + # out-of-bounds lanes zero-fill sBias (cp.async writes 0 on a + # false predicate). K % 8 == 0 with a 32-bit (2 bf16) vector + # means every vector lies wholly in- or out-of-bounds, so one + # predicate per lane is exact. The zero tail is only read back + # for overhang output that the TMA store clamps away anyway. + if epi_tidx < n_active: + bias_pred = cute.make_rmem_tensor( + cute.make_layout((1,)), cutlass.Boolean + ) + bias_pred[0] = cutlass.Boolean( + n_base + epi_tidx * bias_elems_per_copy < mD_mnl.shape[1] + ) + cute.copy_atom_call( + bias_g2s_atom, + gBias_row[(None, epi_tidx)], + sBias_tiled[(None, epi_tidx)], + pred=bias_pred, + ) + cute.arch.cp_async_commit_group() + cute.arch.cp_async_wait_group(0) + self.epilog_sync_barrier.arrive_and_wait() + + # Set tensor memory buffer for current tile + # (T2R, T2R_M, T2R_N, EPI_M, EPI_M) + # Overlapping-accum: the consumer reads the buffer indexed by its + # phase (0/1). Stage-1 acc is strided into TMEM so that its high + # subtiles alias the SFA/SFB columns; to drain those shared columns + # first (before the next tile's MMA overwrites them) the epilogue + # walks subtiles in REVERSE when phase==0. + if cutlass.const_expr(self.overlapping_accum): + acc_stage_index = acc_consumer_state.phase + reverse_subtile = acc_stage_index == 0 + else: + acc_stage_index = acc_consumer_state.index + tTR_tAcc = tTR_tAcc_base[ + (None, None, None, None, None, acc_stage_index) + ] + + # + # Wait for accumulator buffer full + # + acd_pipeline.consumer_wait(acc_consumer_state) + + tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc)) + bSG_gD = cute.group_modes(bSG_gD, 1, cute.rank(bSG_gD)) + if cutlass.const_expr(self.gen_sfd): + # ((T2R, T2R_M, T2R_N), SUBTILE_CNT) + tTR_gSFD = cute.group_modes(tTR_gSFD, 3, cute.rank(tTR_gSFD)) + tTR_cSFD = cute.group_modes(tTR_cSFD, 3, cute.rank(tTR_cSFD)) + + # + # Store accumulator to global memory in subtiles + # + subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3]) + num_prev_subtiles = tile_sched.num_tiles_executed * subtile_cnt + for subtile_idx in cutlass.range(subtile_cnt): + # Map the loop counter to the true output N-subtile. Under + # overlapping-accum with reverse_subtile we walk the output + # subtiles back-to-front so the SFA/SFB-aliased columns are + # drained (and released) before the next tile's MMA reuses + # them. real_subtile_idx addresses every actual output + # position (TMEM acc, gmem D, gmem SFD); the raw subtile_idx + # stays a sequential counter (SMEM ring, release). + real_subtile_idx = subtile_idx + if cutlass.const_expr(self.overlapping_accum): + if reverse_subtile: + real_subtile_idx = ( + self.cta_tile_shape_mnk[1] // self.epi_tile_n + - 1 + - subtile_idx + ) + # + # Load accumulator from tensor memory buffer to register + # + tTR_tAcc_mn = tTR_tAcc[(None, None, None, real_subtile_idx)] + cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc) + + # + # Async arrive accumulator buffer empty earlier when + # overlapping_accum is enabled. Trigger keyed on the raw loop + # counter so it fires exactly once after the shared columns + # (drained first under reverse) have been read out. + # + if cutlass.const_expr(self.overlapping_accum): + if subtile_idx == self.iter_acc_early_release_in_epilogue: + cute.arch.fence_view_async_tmem_load() + with cute.arch.elect_one(): + acd_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + + # Apply per-tensor alpha. + tTR_rAcc.store(tTR_rAcc.load() * alpha) + + # Add per-output-channel bias in FP32 (D = alpha*acc + bias). + # Load this subtile's bias from the CTA-staged sBias row into a + # register fragment shaped like the acc fragment, then add. + # The M-axis stride-0 broadcast means each thread reads only + # its own N-column bias. + if cutlass.const_expr(mBias_mnl is not None): + tTR_rBias = cute.make_rmem_tensor( + tTR_rAcc.shape, mBias_mnl.element_type + ) + cute.copy( + simt_atom_bias, + tTR_sBias[(None, None, None, real_subtile_idx)], + tTR_rBias, + ) + tTR_rAcc.store( + tTR_rAcc.load() + tTR_rBias.load().to(self.acc_dtype) + ) + + # Add the per-element residual in FP32 (D = alpha*acc + bias + # + beta*residual). Wait for this subtile's TMA load to land + # in smem, copy it to registers aligned with the acc + # fragment, then accumulate. The residual shares the output + # dtype and is upconverted to the accumulator type; beta is a + # compile-time constant folded into the scaled add. + if cutlass.const_expr(mC_mnl is not None): + c_pipeline.consumer_wait(c_consumer_state) + cute.copy( + tiled_copy_s2r_c, + tSR_sC[(None, None, None, c_consumer_state.index)], + tSR_rC, + ) + cute.arch.fence_proxy( + "async.shared", + space="cta", + ) + # consumer_release self-elects its signaling threads + # (is_signaling_thread by tidx); all epilogue threads + # call it directly, no explicit elect_one. + c_pipeline.consumer_release(c_consumer_state) + c_consumer_state.advance() + # tSR_rC is the S2R-partition view of tTR_rC + # (shared storage); read back through the acc-shaped view + # so the add lines up with tTR_rAcc. + tTR_rAcc.store( + tTR_rAcc.load() + + self.beta * tTR_rC.load().to(self.acc_dtype) + ) + + # + # SFD generation (NVFP4 output only): per sfd_vec_size + # abs-max -> pvscale_f32 -> cast to sf_dtype (E4M3) -> STG. + # Then rescale acc by norm_const * rcp(qpvscale_f32), where + # qpvscale_f32 is the SFD value read back after the E4M3 cast + # so D_quant * SFD stays self-consistent, with NaN/inf clamp + # via fmin. + # + if cutlass.const_expr(self.gen_sfd): + # Slice gSFD for this subtile and collapse stride-0 broadcast. + gSFD_subtile = tTR_gSFD[(None, None, None, real_subtile_idx)] + t2r_gSFD = cute.filter_zeros(gSFD_subtile) + # The plain SFD STG has no TMA extent clamp, so it must be + # predicated on BOTH the real M and N (channel) extents: + # the cute partition rounds the cta tile up to a full + # mma_tiler multiple in both axes, leaving overhang rows + # (m >= M) and, when cta_n does not divide K, an overhang + # N-subtile (n >= K). An unguarded N-overhang store writes + # zeros at a gmem offset that -- since the mn-row stride + # equals the global sf_k count -- folds onto the next row's + # low sf_k columns, corrupting valid scale factors. The + # fragment lies on one (m, n_base), so the first element's + # coordinates gate the whole store. Index via + # real_subtile_idx: under overlapping-accum reverse_subtile + # the raw loop counter and real output subtile are mirror + # images, so a raw-indexed predicate guards the wrong one. + cSFD_subtile = tTR_cSFD[(None, None, None, real_subtile_idx)] + sfd_m_in_bounds = cute.elem_less( + cSFD_subtile[0][0], mD_mnl.shape[0] + ) + sfd_n_in_bounds = cute.elem_less( + cSFD_subtile[0][1], mD_mnl.shape[1] + ) + sfd_in_bounds = sfd_m_in_bounds and sfd_n_in_bounds + # Partition tTR_rAcc into vec_size groups along the contig (K) mode. + sfgen_rAcc = cute.logical_divide(tTR_rAcc, self.sfd_vec_size) + n_sf = cute.size[1](sfgen_rAcc) + rSFD = cute.make_rmem_tensor((1, n_sf), dtype=self.sfd_dtype) + # Reciprocal of the output dtype's full-scale max (M_D). + rcp_max = cutlass.Float32(1.0 / self.M_D) + fp32_max = cutlass.Float32(3.40282346638528859812e38) + # Single-pass: compute amax, quantize SFD, and rescale acc + # in the same loop. The acc rescale reads the SFD value + # back after the E4M3 cast (qpvscale_f32), so the dequant + # identity D_quant * SFD reproduces the pre-quant + # accumulator to within the FP4 data round-off. + for i_sf in cutlass.range(n_sf, unroll_full=True): + sfgen_slice = sfgen_rAcc[(None, i_sf)] + red_ssa = sfgen_slice.load() + red_abs_ssa = cute.math.absf(red_ssa) + amax = cutlass.Float32( + red_abs_ssa.reduce( + cute.ReductionOp.MAX, + cutlass.Float32(0.0), + 0, + ) + ) + pvscale_f32 = amax * rcp_max * norm_const + rSFD[(0, i_sf)] = pvscale_f32.to(self.sfd_dtype) + qpvscale_f32 = cutlass.Float32( + rSFD[(0, i_sf)].to(cutlass.Float32) + ) + acc_scale = norm_const * cute.arch.rcp_approx(qpvscale_f32) + acc_scale = cutlass.Float32( + nvvm.fmin( + acc_scale.ir_value(), + fp32_max.ir_value(), + nan=True, + ) + ) + sfgen_slice.store(sfgen_slice.load() * acc_scale) + # Store SFD to gmem (predicated on both the M and N + # bounds; the last m-tile and partial-N overhangs have no + # backing SFD storage). Scalar STG per SF element with a + # static (unrolled) index: the gmem address is base + + # i * runtime_stride, so the tensor's global extent may be + # a runtime Int -- only the per-thread element count is + # static. (autovec_copy would require a fully static + # destination and cannot take a runtime-extent tensor.) + if sfd_in_bounds: + for i_st in cutlass.range(n_sf, unroll_full=True): + t2r_gSFD[i_st] = rSFD[(0, i_st)] + + # + # Convert to D type + # + # epilogue_op is applied after the cast, but SFD was computed + # above from the pre-op accumulator -- consistent only for an + # identity op; a non-trivial op would desync D from SFD. + acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load() + acc_vec = epilogue_op(acc_vec.to(self.d_dtype)) + tRS_rD.store(acc_vec) + + # + # Store D to shared memory + # + d_buffer = (num_prev_subtiles + subtile_idx) % self.num_d_stage + cute.copy( + tiled_copy_r2s, + tRS_rD, + tRS_sD[(None, None, None, d_buffer)], + ) + # Fence and barrier to make sure shared memory store is visible to TMA store + cute.arch.fence_proxy( + "async.shared", + space="cta", + ) + self.epilog_sync_barrier.arrive_and_wait() + + # + # TMA store D to global memory + # + if warp_idx == self.epilogue_warp_id[0]: + cute.copy( + tma_atom_d, + bSG_sD[(None, d_buffer)], + bSG_gD[(None, real_subtile_idx)], + ) + # Fence and barrier to make sure shared memory store is visible to TMA store + d_pipeline.producer_commit() + d_pipeline.producer_acquire() + self.epilog_sync_barrier.arrive_and_wait() + + # + # Async arrive accumulator buffer empty. Under overlapping-accum + # the release already happened mid-loop (early release above), so + # only the non-overlapping path releases here. + # + if cutlass.const_expr(not self.overlapping_accum): + with cute.arch.elect_one(): + acd_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + + # + # Advance to next tile (CLC consumer) + # + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + + # + # Dealloc the tensor memory buffer + # + tmem.relinquish_alloc_permit() + self.epilog_sync_barrier.arrive_and_wait() + tmem.free(acc_tmem_ptr) + # + # Wait for D store complete + # + d_pipeline.producer_tail() + + def mainloop_s2t_copy_and_partition( + self, + sSF: cute.Tensor, + tSF: cute.Tensor, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for smem to tmem load for scale factor tensor, then use it to partition smem memory (source) and tensor memory (destination). + + :param sSF: The scale factor tensor in smem + :type sSF: cute.Tensor + :param tSF: The scale factor tensor in tmem + :type tSF: cute.Tensor + + :return: A tuple containing (tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t) where: + - tiled_copy_s2t: The tiled copy operation for smem to tmem load for scale factor tensor(s2t) + - tCsSF_compact_s2t: The partitioned scale factor tensor in smem + - tSF_compact_s2t: The partitioned scale factor tensor in tmem + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + # (MMA, MMA_MN, MMA_K, STAGE) + tCsSF_compact = cute.filter_zeros(sSF) + # (MMA, MMA_MN, MMA_K) + tCtSF_compact = cute.filter_zeros(tSF) + + # Make S2T CopyAtom and tiledCopy + copy_atom_s2t = cute.make_copy_atom( + tcgen05.Cp4x32x128bOp(self.cta_group), + self.sf_dtype, + ) + tiled_copy_s2t = tcgen05.make_s2t_copy(copy_atom_s2t, tCtSF_compact) + thr_copy_s2t = tiled_copy_s2t.get_slice(0) + + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE) + tCsSF_compact_s2t_ = thr_copy_s2t.partition_S(tCsSF_compact) + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K, STAGE) + tCsSF_compact_s2t = tcgen05.get_s2t_smem_desc_tensor( + tiled_copy_s2t, tCsSF_compact_s2t_ + ) + # ((ATOM_V, REST_V), Rest_Tiler, MMA_MN, MMA_K) + tCtSF_compact_s2t = thr_copy_s2t.partition_D(tCtSF_compact) + + return tiled_copy_s2t, tCsSF_compact_s2t, tCtSF_compact_s2t + + def epilog_tmem_copy_and_partition( + self, + tidx: cutlass.Int32, + tAcc: cute.Tensor, + gD_mnl: cute.Tensor, + epi_tile: cute.Tile, + use_2cta_instrs: Union[cutlass.Boolean, bool], + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for tensor memory load, then use it to partition tensor memory (source) and register array (destination). + + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param tAcc: The accumulator tensor to be copied and partitioned + :type tAcc: cute.Tensor + :param gD_mnl: The global tensor D + :type gD_mnl: cute.Tensor + :param epi_tile: The epilogue tiler + :type epi_tile: cute.Tile + :param use_2cta_instrs: Whether use_2cta_instrs is enabled + :type use_2cta_instrs: bool + + :return: A tuple containing (tiled_copy_t2r, tTR_tAcc, tTR_rAcc) where: + - tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + - tTR_tAcc: The partitioned accumulator tensor + - tTR_rAcc: The accumulated tensor in register used to hold t2r results + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + # Make tiledCopy for tensor memory load + copy_atom_t2r = sm100_utils.get_tmem_load_op( + self.cta_tile_shape_mnk, + self.d_layout, + self.d_dtype, + self.acc_dtype, + epi_tile, + use_2cta_instrs, + ) + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, STAGE) + tAcc_epi = cute.flat_divide( + tAcc[((None, None), 0, 0, None)], + epi_tile, + ) + # (EPI_TILE_M, EPI_TILE_N) + tiled_copy_t2r = tcgen05.make_tmem_copy( + copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)] + ) + + thr_copy_t2r = tiled_copy_t2r.get_slice(tidx) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_M, STAGE) + tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi) + + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL) + gD_mnl_epi = cute.flat_divide( + gD_mnl[((None, None), 0, 0, None, None, None)], epi_tile + ) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL) + tTR_gD = thr_copy_t2r.partition_D(gD_mnl_epi) + # (T2R, T2R_M, T2R_N) + tTR_rAcc = cute.make_rmem_tensor( + tTR_gD[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype + ) + return tiled_copy_t2r, tTR_tAcc, tTR_rAcc + + def epilog_smem_load_copy_and_partition( + self, + tiled_copy_t2r: cute.TiledCopy, + tTR_rC: cute.Tensor, + tidx: cutlass.Int32, + sC: cute.Tensor, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """ + Make tiledCopy for shared memory load, then use it to partition register + array (destination) and shared memory (source). Used to read a residual + tile that was TMA-loaded into smem back into registers, aligned to the + accumulator's T2R fragment. + + :param tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + :type tiled_copy_t2r: cute.TiledCopy + :param tTR_rC: The register tensor shaped like the accumulator fragment + :type tTR_rC: cute.Tensor + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param sC: The shared memory tensor to be copied and partitioned + :type sC: cute.Tensor + + :return: A tuple containing (tiled_copy_s2r, tSR_rC, tSR_sC) where: + - tiled_copy_s2r: The tiled copy operation for smem to register copy(s2r) + - tSR_rC: The partitioned register tensor (destination) + - tSR_sC: The partitioned shared memory tensor (source) + :rtype: Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor] + """ + copy_atom_s2r = cute.make_copy_atom(cute.nvgpu.CopyUniversalOp(), self.d_dtype) + tiled_copy_s2r = cute.make_tiled_copy_D(copy_atom_s2r, tiled_copy_t2r) + # (S2R, S2R_M, S2R_N, PIPE) + thr_copy_s2r = tiled_copy_s2r.get_slice(tidx) + tSR_sC = thr_copy_s2r.partition_D(sC) + # (S2R, S2R_M, S2R_N) + tSR_rC = tiled_copy_s2r.retile(tTR_rC) + return tiled_copy_s2r, tSR_rC, tSR_sC + + def epilog_gmem_copy_and_partition( + self, + tidx: cutlass.Int32, + atom: Union[cute.CopyAtom, cute.TiledCopy], + gD_mnl: cute.Tensor, + epi_tile: cute.Tile, + sD: cute.Tensor, + ) -> Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor]: + """Make tiledCopy for global memory store, then use it to: + partition shared memory (source) and global memory (destination) for TMA store version. + + :param tidx: The thread index in epilogue warp groups + :type tidx: cutlass.Int32 + :param atom: The copy_atom_d to be used for TMA store version, or tiled_copy_t2r for none TMA store version + :type atom: cute.CopyAtom or cute.TiledCopy + :param gD_mnl: The global tensor D + :type gD_mnl: cute.Tensor + :param epi_tile: The epilogue tiler + :type epi_tile: cute.Tile + :param sD: The shared memory tensor to be copied and partitioned + :type sD: cute.Tensor + + :return: A tuple containing (tma_atom_d, bSG_sD, bSG_gD) where: + - tma_atom_d: The TMA copy atom + - bSG_sD: The partitioned shared memory tensor D + - bSG_gD: The partitioned global tensor D + :rtype: Tuple[cute.CopyAtom, cute.Tensor, cute.Tensor] + """ + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL) + gD_epi = cute.flat_divide( + gD_mnl[((None, None), 0, 0, None, None, None)], epi_tile + ) + + tma_atom_d = atom + sD_for_tma_partition = cute.group_modes(sD, 0, 2) + gD_for_tma_partition = cute.group_modes(gD_epi, 0, 2) + # ((ATOM_V, REST_V), EPI_M, EPI_N) + # ((ATOM_V, REST_V), EPI_M, EPI_N, RestM, RestN, RestL) + bSG_sD, bSG_gD = cpasync.tma_partition( + tma_atom_d, + 0, + cute.make_layout(1), + sD_for_tma_partition, + gD_for_tma_partition, + ) + return tma_atom_d, bSG_sD, bSG_gD + + @staticmethod + def _make_sf_smem_layouts( + tiled_mma: cute.TiledMma, + flat_mma_tiler_sf: Tuple[int, int, int], + sf_vec_size: int, + ab_stage: int, + ) -> Tuple[cute.Layout, cute.Layout]: + """Build the staged shared memory layouts for the A and B scale factors. + + These take the flat-K SF tiler: the blockscaled utils expect a scalar K + where conv carries a (K,) tuple, and K runs on the SF cadence. + + :param tiled_mma: The tiled MMA the scale factors feed + :param flat_mma_tiler_sf: The MMA tiler with a flat K on the SF cadence + :param sf_vec_size: Number of channels one scale factor covers + :param ab_stage: Number of A/B pipeline stages + + :return: The SFA and SFB layouts, each staged + """ + sfa_smem_layout_staged = blockscaled_utils.make_smem_layout_sfa( + tiled_mma, flat_mma_tiler_sf, sf_vec_size, ab_stage + ) + sfb_smem_layout_staged = blockscaled_utils.make_smem_layout_sfb( + tiled_mma, flat_mma_tiler_sf, sf_vec_size, ab_stage + ) + return sfa_smem_layout_staged, sfb_smem_layout_staged + + def _make_acc_fake_tensor(self, tiled_mma, mma_tiler): + """Build the accumulator fake tensor used for TMEM column accounting. + + For overlapping-accum the second logical acc buffer is strided into the + columns otherwise reserved for SFA/SFB (so the next tile's math can + overlap the current tile's epilogue). The stride for the stage dim is + (cta_tile_n - num_sf_tmem_cols) * stride[0][1], the block-scaled GEMM + overlapping-accum layout. Otherwise fall back to a plain num_acc_stage fake. + """ + acc_shape = tiled_mma.partition_shape_C(mma_tiler[:2]) + if cutlass.const_expr(self.overlapping_accum): + tCtAcc_fake = tiled_mma.make_fragment_C(cute.append(acc_shape, 2)) + s = tCtAcc_fake.stride + return cute.make_tensor( + tCtAcc_fake.iterator, + cute.make_layout( + tCtAcc_fake.shape, + stride=( + s[0], + s[1], + s[2], + (self.cta_tile_shape_mnk[1] - self.num_sf_tmem_cols) * s[0][1], + ), + ), + ) + return tiled_mma.make_fragment_C(cute.append(acc_shape, self.num_acc_stage)) + + def _compute_stages( + self, + tiled_mma: cute.TiledMma, + mma_tiler_mnk: Tuple[int, int, int], + a_dtype: Type[cutlass.Numeric], + b_dtype: Type[cutlass.Numeric], + epi_tile: cute.Tile, + d_dtype: Type[cutlass.Numeric], + d_layout: LayoutEnum, + sf_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + flat_mma_tiler_sf: Tuple[int, int, int], + smem_capacity: int, + occupancy: int, + has_residual: bool = False, + ) -> Tuple[int, int, int]: + """Computes the number of stages for A/B/D operands. + + The A/B stage count is chosen so the full SharedStorage (every pipeline + mbar, whose count scales with the stage count, plus every operand array + with its 1024B alignment padding) fits in smem. The number is derived by + constructing the actual SharedStorage struct for a candidate stage count + and reading its exact size_in_bytes(), then taking the largest count that + fits -- rather than dividing capacity by a per-stage byte estimate that + omits the alignment padding and the stage-scaled mbar bytes. + + :param tiled_mma: The tiled MMA object defining the core computation. + :param mma_tiler_mnk: The shape (M, N, K) of the MMA tiler. + :param a_dtype: Data type of operand A. + :param b_dtype: Data type of operand B. + :param epi_tile: The epilogue tile shape. + :param d_dtype: Data type of operand D (output). + :param d_layout: Layout enum of operand D. + :param sf_dtype: Data type of Scale factor. + :param sf_vec_size: Scale factor vector size. + :param flat_mma_tiler_sf: The flat-K mma tiler the SF layouts are built + from: the A/B tiler's M and N with the SF cadence on K. + :param smem_capacity: Total available shared memory capacity in bytes. + :param occupancy: CTAs per SM. Always 1 for this persistent kernel. + + :return: (ACC stages, A/B operand stages, D stages) + :rtype: tuple[int, int, int] + """ + # ACC stages + num_acc_stage = 1 if mma_tiler_mnk[1] == 256 else 2 + + # D stages + num_d_stage = 2 + + buffer_align_bytes = 1024 + num_clc_stage = 1 + num_clc_response_bytes = 16 + cta_tile_n = mma_tiler_mnk[1] + + def smem_bytes(ab_stage: int, d_stage: int) -> int: + """Exact SharedStorage byte size for a candidate stage count. + + Builds the operand/SF/epi layouts for this ``ab_stage``/``d_stage`` + and assembles the SharedStorage this kernel will allocate. Every field + mirrors the runtime SharedStorage struct defined in __call__; keep the + two in sync so the byte total is exact. + """ + a_smem, b_smem, d_smem = self._make_smem_layouts( + tiled_mma, + mma_tiler_mnk, + a_dtype, + b_dtype, + d_dtype, + d_layout, + epi_tile, + ab_stage, + d_stage, + ) + sfa_smem, sfb_smem = self._make_sf_smem_layouts( + tiled_mma, flat_mma_tiler_sf, sf_vec_size, ab_stage + ) + + @cute.struct + class ProbeStorage: + """Field-for-field mirror of the runtime SharedStorage, sized for + this stage count so size_in_bytes() gives the exact allocation.""" + + ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, ab_stage * 2] + sfa_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, ab_stage * 2] + acc_full_mbar_ptr: cute.struct.MemRange[ + cutlass.Int64, num_acc_stage * 2 + ] + tmem_dealloc_mbar_ptr: cutlass.Int64 + tmem_holding_buf: cutlass.Int32 + clc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_clc_stage * 2] + clc_response: cute.struct.Align[ + cute.struct.MemRange[ + cutlass.Int32, num_clc_response_bytes // 4 * num_clc_stage + ], + 16, + ] + sD: cute.struct.Align[ + cute.struct.MemRange[d_dtype, cute.cosize(d_smem.outer)], + buffer_align_bytes, + ] + sBias: cute.struct.Align[ + cute.struct.MemRange[d_dtype, cta_tile_n if self.has_bias else 0], + buffer_align_bytes, + ] + sC: cute.struct.Align[ + cute.struct.MemRange[ + d_dtype, + cute.cosize(d_smem.outer) if has_residual else 0, + ], + buffer_align_bytes, + ] + c_full_mbar_ptr: cute.struct.MemRange[ + cutlass.Int64, d_stage * 2 if has_residual else 0 + ] + sA: cute.struct.Align[ + cute.struct.MemRange[a_dtype, cute.cosize(a_smem.outer)], + buffer_align_bytes, + ] + sB: cute.struct.Align[ + cute.struct.MemRange[b_dtype, cute.cosize(b_smem.outer)], + buffer_align_bytes, + ] + sSFA: cute.struct.Align[ + cute.struct.MemRange[sf_dtype, cute.cosize(sfa_smem)], + buffer_align_bytes, + ] + sSFB: cute.struct.Align[ + cute.struct.MemRange[sf_dtype, cute.cosize(sfb_smem)], + buffer_align_bytes, + ] + + return ProbeStorage.size_in_bytes() + + # A single A/B stage must fit: a tile whose one-stage SharedStorage already + # exceeds capacity is unimplementable and is rejected upstream by + # can_implement, so the lower-bound search below can assume >= 1 fits. + if smem_bytes(1, num_d_stage) > smem_capacity: + raise RuntimeError( + "tile too large for even one A/B stage; can_implement should have " + "rejected it before reaching stage selection" + ) + + # Compute a lower bound that is guaranteed to fit, then scan upward once + # (never downward) to the largest A/B stage whose exact SharedStorage + # fits. Each added stage grows the four staged arrays (sA/sB/sSFA/sSFB) by + # their raw per-stage bytes; the Align[1024] rounding of each array is a + # one-time boundary crossing over the whole growth, so it contributes at + # most 4*1024 B total. Bounding the alignment loss as that constant in the + # numerator (rather than inflating every stage's slope) keeps the lower + # bound within a few stages of the true maximum while still guaranteeing + # cost(lb) <= cost(1) + (lb-1)*slope + 4*1024 <= capacity. A/B depth is + # filled first (it dominates mainloop overlap); any smem left over then + # deepens the epilogue. size_in_bytes is monotonic in each stage count, so + # the upward passes land on the true maximum. + a_one, b_one, _ = self._make_smem_layouts( + tiled_mma, + mma_tiler_mnk, + a_dtype, + b_dtype, + d_dtype, + d_layout, + epi_tile, + 1, + 1, + ) + sfa_one, sfb_one = self._make_sf_smem_layouts( + tiled_mma, flat_mma_tiler_sf, sf_vec_size, 1 + ) + slope = ( + cute.cosize(a_one.outer) * a_dtype.width // 8 + + cute.cosize(b_one.outer) * b_dtype.width // 8 + + cute.cosize(sfa_one) * sf_dtype.width // 8 + + cute.cosize(sfb_one) * sf_dtype.width // 8 + + 2 * 8 # ab_full mbar (Int64) + + 2 * 8 # sfa_full mbar (Int64) + ) + align_loss = 4 * buffer_align_bytes + num_ab_stage = max( + 1, 1 + (smem_capacity - smem_bytes(1, num_d_stage) - align_loss) // slope + ) + while smem_bytes(num_ab_stage + 1, num_d_stage) <= smem_capacity: + num_ab_stage += 1 + while smem_bytes(num_ab_stage, num_d_stage + 1) <= smem_capacity: + num_d_stage += 1 + + return num_acc_stage, num_ab_stage, num_d_stage + + def can_implement( + self, + c: int, + k: int, + ab_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], + sf_dtype: Type[cutlass.Numeric], + output: cute.Tensor, + filter_trs: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + dil_dhw: Tuple[int, int, int], + c_dtype: Optional[Type[cutlass.Numeric]] = None, + has_bias: bool = False, + epilogue_op: Optional[cutlass.Constexpr] = None, + ) -> bool: + """Determine if the given tensor configuration can be implemented by this kernel.""" + try: + # Residual (C) reuses the output D im2col TMA descriptor, so it must + # match D's dtype exactly. c_dtype is None when no residual is passed. + if c_dtype is not None and c_dtype is not d_dtype: + raise testing.CantImplementError( + f"Residual c_dtype ({c_dtype}) must equal d_dtype ({d_dtype})." + ) + # Two input formats, each pinned to one scale-factor format: + # NVFP4 : Float4E2M1FN A/B, Float8E4M3FN scale over 16 channels. + # MXFP8 : Float8E4M3FN A/B, Float8E8M0FNU scale over 32 channels. + # The element type is pinned alongside the vector size because it drives + # the SFB SMEM layout, the barrier transaction bytes and the SFD rounding + # -- a mismatched one is misinterpreted on the device, not refused. SFD + # takes the input's format, so this pins both sides. + sf_format = { + cutlass.Float4E2M1FN: (cutlass.Float8E4M3FN, 16), + cutlass.Float8E4M3FN: (cutlass.Float8E8M0FNU, 32), + }.get(ab_dtype) + if sf_format is None: + raise testing.CantImplementError( + f"Only Float4E2M1FN (NVFP4) and Float8E4M3FN (MXFP8) A/B " + f"are supported, got {ab_dtype}." + ) + want_sf_dtype, want_sf_vec_size = sf_format + if sf_dtype is not want_sf_dtype or self.sf_vec_size != want_sf_vec_size: + raise testing.CantImplementError( + f"{ab_dtype} A/B requires sf_dtype={want_sf_dtype} over " + f"{want_sf_vec_size} channels, got sf_dtype={sf_dtype}, " + f"sf_vec_size={self.sf_vec_size}." + ) + # D is either the narrow block-scaled output that pairs with the input + # format -- NVFP4 with FP4, MXFP8 with FP8 E4M3 -- or a 16-bit wide output. + # The pairing is forced by the SFD: the epilogue rescales the accumulator by + # an approximate reciprocal of the quantized scale factor, exact for E8M0's + # powers of two but not for E4M3's arbitrary values. FP4's eight magnitudes + # absorb that error; FP8 E4M3's few hundred do not, leaving a data-dependent + # fraction of elements one step off. A wide output carries no SFD and stops + # at 16 bits: FP32 D doubles the epilogue buffer and takes A/B stages with + # it, and what reads this kernel's output is narrow or 16-bit. + paired_narrow_output = { + cutlass.Float4E2M1FN: cutlass.Float4E2M1FN, + cutlass.Float8E4M3FN: cutlass.Float8E4M3FN, + }[ab_dtype] + if d_dtype not in ( + paired_narrow_output, + cutlass.BFloat16, + cutlass.Float16, + ): + raise testing.CantImplementError( + f"d_dtype must be {paired_narrow_output} (the block-scaled narrow " + f"output paired with {ab_dtype} A/B), BFloat16 or Float16, got " + f"{d_dtype}." + ) + # A block-scaled output constrains the whole epilogue. bias and residual + # take D's own dtype, so the addend would be FP4 or FP8 -- eight magnitudes + # in one case, a few hundred in the other -- making the term's quantization + # error the same order as the term itself. And the SFD is taken from the + # accumulator before the epilogue op runs, so anything but identity leaves D + # scaled by a factor that does not describe it. The op is matched against + # the named non-identity activations: a caller's own pass-through lambda is + # identity in effect and must not be refused for being a different object. + if d_dtype in (cutlass.Float4E2M1FN, cutlass.Float8E4M3FN): + if has_bias: + raise testing.CantImplementError( + f"A block-scaled {d_dtype} output does not support a bias." + ) + if c_dtype is not None: + raise testing.CantImplementError( + f"A block-scaled {d_dtype} output does not support a residual." + ) + if epilogue_op in { + entry["device"] + for name, entry in EPILOGUE_ACTIVATIONS.items() + if name != "identity" + }: + raise testing.CantImplementError( + f"A block-scaled {d_dtype} output only supports the identity " + f"epilogue op." + ) + # One epilogue lane stages one 32-bit vector of the bias row, so the row + # takes cta_tile_n / (32 / d_dtype.width) lanes and the whole row is filled + # in a single pass. Every admitted output is 16 bits or narrower, which + # keeps the widest N tile inside the lane count; stating the bound makes a + # wider tile or a wider output refuse here rather than read back shared + # memory that no lane wrote. + if has_bias: + bias_lanes = self.mma_tiler_mn[1] // (32 // d_dtype.width) + epilogue_lanes = 32 * len(self.epilogue_warp_id) + if bias_lanes > epilogue_lanes: + raise testing.CantImplementError( + f"Staging a {d_dtype} bias row for a {self.mma_tiler_mn[1]}-wide " + f"N tile takes {bias_lanes} lanes, more than the " + f"{epilogue_lanes} the epilogue has." + ) + # The kernel emits the NZPQK output with K (the GEMM-N axis) contiguous, + # i.e. n-major. The SFD store and TMA epilogue assume this; any other + # leading dim would write the wrong axis. Reject it up front. + if output.leading_dim != 4: + raise testing.CantImplementError( + "D must be n-major (NZPQK with K contiguous, leading_dim=4), " + f"got leading_dim={output.leading_dim}" + ) + self.check_mma_tiler_and_cluster_shape() + # Implicit-GEMM M is N*Z*P*Q, the D tensor's leading four modes. + n_out, z_out, p_out, q_out = output.shape[:4] + self.check_preferred_cluster_problem_size(n_out * z_out * p_out * q_out, k) + # Blockscaled MMA requires a per-CTA M of 128: make_blockscaled_trivial_tiled_mma + # has no M=64 atom and raises "expects the M-mode to be 128, but got 64" at + # cute.compile time. The per-CTA M is mma_tiler_m halved under 2CTA, so both + # mma_tiler_m=64 (1CTA) and mma_tiler_m=128 (2CTA) hit M=64 and fail. The + # parent check_mma_tiler_and_cluster_shape admits both (it allows 64/128 for + # 1CTA and 128/256 for 2CTA from the dense-GEMM contract), so reject them here + # with a clear message instead of the opaque MMA-builder error. The two viable + # shapes are mma_tiler_m=128 (1CTA) and mma_tiler_m=256 (2CTA), both giving M=128. + per_cta_m = self.mma_tiler_mn[0] // (2 if self.use_2cta_instrs else 1) + if per_cta_m != 128: + raise testing.CantImplementError( + f"Blockscaled conv requires per-CTA M=128, got {per_cta_m} from " + f"mma_tiler_m={self.mma_tiler_mn[0]}, use_2cta_instrs=" + f"{self.use_2cta_instrs}. Use mma_tiler_m=128 with 1CTA or " + f"mma_tiler_m=256 with 2CTA." + ) + # SFB stages 128 output channels at a time, so an mma_tiler_n below that + # puts several N tiles in one chunk, each needing its own Tensor Memory + # read offset. That offset exists for 64 (two tiles per chunk) and 192 + # (the chunk pair reshaped into overlapping 256-wide windows); 128 and 256 + # sit on chunk boundaries. The rest silently read the chunk's first tile + # once there is more than one N tile, so they serve a single one. + cta_tile_n = self.mma_tiler_mn[1] + if cta_tile_n not in (64, 128, 192, 256): + n_tiles = -(-k // cta_tile_n) + if n_tiles > 1: + raise testing.CantImplementError( + f"mma_tiler_n={cta_tile_n} has no per-tile SFB chunk offset, " + f"so it only serves a single N tile, but K={k} needs " + f"{n_tiles}. Use mma_tiler_n of 64, 128, 192 or 256." + ) + # 16B output alignment forces Kout % 32 == 0 for FP4 D, which is + # what makes the SFD N-bound predicate (gates a whole subtile on its + # first-element N coord) exact -- so no separate Kout guard is needed. + _check_tensor_alignment(c, k, ab_dtype, d_dtype) + # C legality. Either format takes a partial trailing K tile -- its A and B + # arrive zero-filled from the TMA and its scale factors are read in bounds -- + # so C need not be a whole number of K tiles. What it must be is a whole + # number of scale-factor atoms, 4 * sf_vec_size channels: SFA's 4-byte + # cp.async carries exactly one atom and walks a row stride of C / + # sf_vec_size factors, which that alignment keeps a multiple of the 4 the + # transfer needs. A C that splits an atom would leave the row stride + # unaligned, and padding the host allocation to hide it would put the + # factors a position owns at an offset the kernel does not address. + sf_atom_channels = self.sf_vec_size * 4 + if c % sf_atom_channels != 0: + raise testing.CantImplementError( + f"{ab_dtype} requires C to be a multiple of {sf_atom_channels}, " + f"got C = {c}." + ) + # cta_tile_k legality. Both formats need a whole number of MMA K + # instructions: mma_inst_tile_k truncates the division, so a tile_k off + # that multiple silently runs a narrower tile than asked for. 192 is a + # legal multiple but excluded: its SF-K atom count is three, and an odd + # count above one cannot be halved when the SFB multicast splits along K + # under 2CTA, which deadlocks the tmem-alloc mbarrier. A K tile need not + # divide C, so nothing needs it. + # MXFP8 stops at 128, the width of its SF atom. Below it the atom is + # shared: a 64-channel K tile stages the whole 128-channel atom and reads + # half of it, which only works because the two halves divide it exactly. + # Above it a K tile would span several atoms per tile, which no + # configuration here exercises. + # NVFP4's atom is 64 channels, at or below every legal tile, so its upper + # end is set by validation instead: 256 is four MMA K instructions, wider + # tiles are unvalidated and tracing time climbs past them -- 768 does not + # finish in 400 s. + allowed_tile_k = ( + (64, 128) if ab_dtype is cutlass.Float8E4M3FN else (64, 128, 256) + ) + if self.cta_tile_k not in allowed_tile_k: + raise testing.CantImplementError( + f"{ab_dtype} requires cta_tile_k in {allowed_tile_k}, got " + f"{self.cta_tile_k}." + ) + # The SFB multicast hands each cluster-M CTA a part of the staged scale + # factors, splitting them along K, so the cluster cannot be deeper in M + # than the K tile has scale factors to give out. The tile_k bound above + # only settles whether an atom count can be halved, which is the two CTAs + # one MMA atom spans; a cluster wider than that asks for more parts. Past + # the limit the descriptor addresses a part that is not there and the + # store faults on a misaligned address, so it is refused up front. MXFP8 + # stages one atom whatever the tile_k, capping it at four either way. + sf_per_k_tile = ( + sf_k_tile_channels(self.cta_tile_k, self.sf_vec_size) + // self.sf_vec_size + ) + for which, cluster_mn in ( + ("preferred", self.preferred_cluster_shape_mn), + ("fallback", self.fallback_cluster_shape_mn), + ): + if cluster_mn[0] > sf_per_k_tile: + raise testing.CantImplementError( + f"{ab_dtype} with cta_tile_k = {self.cta_tile_k} stages " + f"{sf_per_k_tile} scale factors per K tile, so the {which} " + f"cluster cannot span more than {sf_per_k_tile} CTAs along M, " + f"got {cluster_mn[0]}." + ) + _check_im2col_descriptor_limits( + filter_trs, + upper_padding_dhw, + lower_padding_dhw, + stride_dhw, + dil_dhw, + ) + # The inherited cluster checks raise cutlass.testing.CantImplementError, + # which every local raise derives from, so catch it at that width. + except testing.CantImplementError as e: + print(e) + return False + return True + + +def sf_k_tile_channels(cta_tile_k: int, sf_vec_size: int) -> int: + """Channels one SF K tile spans. + + The SF atom holds 4 scale-factor blocks along K, so it covers + 4 * sf_vec_size channels and cannot be staged in part: a narrower A/B K tile + still stages a whole atom, and several A/B tiles then share one SF tile. Host + and kernel both derive the SF cadence from here so they cannot drift apart. + """ + return max(cta_tile_k, 4 * sf_vec_size) + + +def sfb_per_position_channels(c: int, cta_tile_k: int, sf_vec_size: int) -> int: + """Channels of SFB one filter position spans: C rounded up to whole SF K tiles. + + SFB runs channels and filter positions together in one flat K mode, so a K tile + that does not divide C would straddle two positions; rounding each position out + to whole SF K tiles puts every position on a tile boundary. The bound is the SF + tile, the width of the TMA box: a C that a 64-channel A/B tile divides still + straddles two positions once a 128-channel SF chunk serves two of those tiles. + """ + span = sf_k_tile_channels(cta_tile_k, sf_vec_size) + return -(-c // span) * span + + +# Exempt from the frontend cache-key rule (see indexer/mqa_logits_util.py +# frontend_cache_token): the key below contains cute.Tensor arguments, which +# hash by identity and are rebuilt via from_dlpack on every run(), so this +# cache can never hit across calls, let alone across frontends. If the Tensor +# arguments are ever dropped from the key or Tensor gains structural hashing, +# add the frontend token here. +@lru_cache(maxsize=1) +def compile_conv( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + input: cute.Tensor, + filter: cute.Tensor, + output: cute.Tensor, + acc_dtype: Type[cutlass.Numeric], + sfa: cute.Tensor, + sfb: cute.Tensor, + sfd: Optional[cute.Tensor], + alpha: cutlass.Float32, + norm_const: cutlass.Float32, + bias: Optional[cute.Tensor], + residual: Optional[cute.Tensor], + beta: float, + sf_vec_size: int, + cta_tile_k: int, + mma_tiler: Tuple[int, int] = (256, 256), + preferred_cluster_shape_mn: Tuple[int, int] = (2, 1), + fallback_cluster_shape_mn: Tuple[int, int] = (2, 1), + swizzle_size: int = 1, + raster_along: Literal["m", "n"] = "m", + use_2cta_instrs: bool = True, + upper_padding_dhw: Tuple[int, int, int] = (0, 0, 0), + lower_padding_dhw: Tuple[int, int, int] = (0, 0, 0), + stride_dhw: Tuple[int, int, int] = (1, 1, 1), + dilation_dhw: Tuple[int, int, int] = (1, 1, 1), + epilogue_op: cutlass.Constexpr = lambda x: x, +): + """ + Compile a 3D convolution kernel. + + :param ncdhw: Problem shape (N, C, D, H, W) + :param ktrs: Problem shape (K, T, R, S) + :param input: Input tensor (N, D, H, W, C) with C contiguous + :param filter: Filter tensor in KTRSC format (K, T, R, S, C) with C contiguous + :param output: Output tensor (N, Z, P, Q, K) with K contiguous + :param acc_dtype: Accumulator data type + :param mma_tiler: MMA tile shape (M, N) + :param preferred_cluster_shape_mn: Preferred cluster shape (M, N) for CLC dynamic scheduling + :param fallback_cluster_shape_mn: Fallback cluster shape (M, N) for CLC dynamic scheduling + :param swizzle_size: Swizzling size in the unit of cluster. 1 means no swizzle + :param raster_along: Rasterization order of clusters. Only used when swizzle_size > 1 + :param use_2cta_instrs: Whether to use 2CTA instructions + :param upper_padding_dhw: Upper padding (PadD, PadH, PadW) + :param lower_padding_dhw: Lower padding (PadD, PadH, PadW) + :param stride_dhw: Stride (Sd, Sh, Sw) + :param dilation_dhw: Dilation (DilD, DilH, DilW) + :param epilogue_op: Epilogue operation + :return: Compiled kernel function + """ + from cutlass.cute.runtime import make_fake_stream + + # Output spatial dims for the SFD global descriptor (host int, trace-const). + zpq = compute_zpq( + ncdhw[2:], + ktrs[1:], + stride_dhw, + upper_padding_dhw, + lower_padding_dhw, + dilation_dhw, + ) + # Host-side swizzle_size guard: reject a swizzle that exceeds the cluster + # count in the swizzled dimension (GEMM-M = N*Z*P*Q spatial, GEMM-N = Kout). + _check_swizzle_size( + ncdhw[0] * zpq[0] * zpq[1] * zpq[2], + ktrs[0], + mma_tiler, + use_2cta_instrs, + preferred_cluster_shape_mn, + fallback_cluster_shape_mn, + swizzle_size, + raster_along, + ) + + # Create convolution kernel object. Only cta_tile_k (compile-time K tile) is + # a build-time shape parameter; the input channel count C and all output + # geometry (N/Z/P/Q/K) are read from the dynamic tensors at runtime, so one + # cubin serves every C. Filter T/R/S and pad/stride/dil are likewise runtime + # (dynamic filter extents + boxed Int32). + conv_op = Sm100BlockScaledPersistentDenseImplicitGemmFpropKernel( + acc_dtype, + sf_vec_size, + use_2cta_instrs, + mma_tiler, + preferred_cluster_shape_mn, + fallback_cluster_shape_mn, + cta_tile_k, + swizzle_size, + raster_along, + ) + + # Check if configuration can be implemented + can_implement = conv_op.can_implement( + ncdhw[1], + ktrs[0], + input.element_type, + output.element_type, + sfa.element_type, + output, + ktrs[1:], + upper_padding_dhw, + lower_padding_dhw, + stride_dhw, + dilation_dhw, + residual.element_type if residual is not None else None, + bias is not None, + epilogue_op, + ) + if not can_implement: + raise testing.CantImplementError("The current config is invalid/unsupported.") + + stream = make_fake_stream() + # Box pad/stride/dilation as cutlass.Int32 so the cute.compile entry scalars + # lower to runtime SSA AND keep their values out of the mangled function + # name; that combination is what lets one cubin be reused across configs. A + # raw Python int here would fold the value into the name and force + # recompilation per config. The filter T/R/S are already runtime (taken from + # the dynamic-layout filter tensor extents). + return cute.compile( + conv_op, + input, + filter, + output, + sfa, + sfb, + alpha, + epilogue_op, + sfd, + norm_const, + bias, + residual, + beta, + *_rt_conv_scalars( + upper_pad=upper_padding_dhw, + lower_pad=lower_padding_dhw, + stride=stride_dhw, + dil=dilation_dhw, + ), + stream, + ) + + +def create_cute_tensor( + source_f32_tensor: torch.Tensor, + dtype: Type[cutlass.Numeric], + leading_dim: int = None, +) -> Tuple[cute.Tensor, torch.Tensor]: + """Create a dynamic-layout cute tensor from a source f32 tensor. + + The tensor is always marked dynamic-layout: its non-leading extents lower to + runtime SSA so one compiled cubin serves any N/C/D/H/W/K/T/R/S config. + + :param source_f32_tensor: Source f32 tensor + :type source_f32_tensor: torch.Tensor + :param dtype: Data type + :type dtype: Type[cutlass.Numeric] + :param leading_dim: Leading dimension kept contiguous in the dynamic layout + :type leading_dim: int + :return: Tuple of cute tensor and storage tensor + :rtype: Tuple[cute.Tensor, torch.Tensor] + """ + + # FP4 needs packed storage: two elements per byte. Build the buffer here, with + # the trailing dim halved, rather than through the shared cute_tensor_like, + # which sizes a uint8 buffer at one byte per source element. + if dtype == cutlass.Float4E2M1FN: + shape = tuple(source_f32_tensor.shape) + assert shape[-1] % 2 == 0, ( + f"FP4 packed storage requires trailing dim even, got shape={shape}" + ) + packed_shape = shape[:-1] + (shape[-1] // 2,) + storage_int8 = torch.empty(packed_shape, dtype=torch.int8, device="cuda") + storage_view = storage_int8.view(dtype=torch.float4_e2m1fn_x2) + cute_tensor = from_dlpack(storage_view, assumed_align=16) + cute_tensor = cute_tensor.mark_layout_dynamic(leading_dim=leading_dim) + + if source_f32_tensor.numel() > 0: + f32_gpu = ( + source_f32_tensor.cuda() + if source_f32_tensor.device.type == "cpu" + else source_f32_tensor + ) + f32_cute = from_dlpack(f32_gpu) + f32_cute = f32_cute.mark_layout_dynamic(leading_dim=leading_dim) + cute.testing.convert(f32_cute, cute_tensor) + return cute_tensor, storage_view + + cute_tensor, storage_tensor = cutlass_torch.cute_tensor_like( + source_f32_tensor, dtype, is_dynamic_layout=True, assumed_align=16 + ) + + return cute_tensor, storage_tensor + + +@cute.jit +def cvt_sf_MKL_to_M32x4xrm_K4xrk_L( + sf_ref_tensor: cute.Tensor, + sf_mma_tensor: cute.Tensor, +): + """Convert scale factor tensor from MKL layout to mma specification M(32x4xrest_m)xK(4xrest_k)xL layout""" + # sf_mma_tensor has flatten shape (32, 4, rest_m, 4, rest_k, l) + # group to ((32, 4, rest_m), (4, rest_k), l) + sf_mma_tensor = cute.group_modes(sf_mma_tensor, 0, 3) + sf_mma_tensor = cute.group_modes(sf_mma_tensor, 1, 3) + for i in cutlass.range(cute.size(sf_ref_tensor)): + mkl_coord = sf_ref_tensor.layout.get_hier_coord(i) + sf_mma_tensor[mkl_coord] = sf_ref_tensor[mkl_coord] + + +# Compile-time epilogue activations, folded into the cubin as a Constexpr op +# (one activation per cubin). Each entry pairs the device-side op applied to the +# output fragment with the torch op used to build the reference. +EPILOGUE_ACTIVATIONS = { + "identity": { + "device": lambda x: x, + "ref": lambda x: x, + }, + "relu": { + "device": lambda x: cute.where(x > 0, x, cute.full_like(x, 0)), + "ref": torch.nn.functional.relu, + }, +} + + +# Create scale factor tensor SFA/SFB +def create_scale_factor_tensor_swizzled(L, mn, k, sf_vec_size, dtype): + def ceil_div(a, b): + return (a + b - 1) // b + + sf_k = ceil_div(k, sf_vec_size) + ref_shape = (L, mn, sf_k) + + atom_m = (32, 4) + atom_k = 4 + mma_shape = ( + L, + ceil_div(mn, atom_m[0] * atom_m[1]), + ceil_div(sf_k, atom_k), + atom_m[0], + atom_m[1], + atom_k, + ) + + ref_permute_order = (1, 2, 0) + mma_permute_order = (3, 4, 1, 5, 2, 0) + + # Create f32 ref torch tensor (cpu) + ref_f32_torch_tensor_cpu = cutlass_torch.create_and_permute_torch_tensor( + ref_shape, + torch.float32, + permute_order=ref_permute_order, + init_type=cutlass_torch.TensorInitType.RANDOM, + init_config=cutlass_torch.RandomInitConfig( + min_val=1, + max_val=3, + ), + ) + # Create f32 cute torch tensor (cpu) + cute_f32_torch_tensor_cpu = cutlass_torch.create_and_permute_torch_tensor( + mma_shape, + torch.float32, + permute_order=mma_permute_order, + init_type=cutlass_torch.TensorInitType.RANDOM, + init_config=cutlass_torch.RandomInitConfig( + min_val=0, + max_val=1, + ), + ) + + # convert ref f32 tensor to cute f32 tensor + cvt_sf_MKL_to_M32x4xrm_K4xrk_L( + from_dlpack(ref_f32_torch_tensor_cpu), + from_dlpack(cute_f32_torch_tensor_cpu), + ) + cute_f32_torch_tensor = cute_f32_torch_tensor_cpu.cuda() + + # reshape makes memory contiguous + ref_f32_torch_tensor_cpu = ( + ref_f32_torch_tensor_cpu.permute(2, 0, 1) + .unsqueeze(-1) + .expand(L, mn, sf_k, sf_vec_size) + .reshape(L, mn, sf_k * sf_vec_size) + .permute(*ref_permute_order) + ) + # prune to mkl for reference check. + ref_f32_torch_tensor_cpu = ref_f32_torch_tensor_cpu[:, :k, :] + + # Round-trip the reference scale factors through the storage dtype so the + # reference dequant uses the exact values the kernel reads. E4M3 keeps + # fractional values (near-lossless for the 1..3 init range), but E8M0 + # (MXFP8) snaps every scale to a power of two, so an un-rounded f32 + # reference would disagree with the kernel on every non-pow2 block. + ref_f32_torch_tensor_cpu = ref_f32_torch_tensor_cpu.to( + cutlass_torch.dtype(dtype) + ).to(torch.float32) + + # Create dtype cute torch tensor (cpu) + cute_tensor, cute_torch_tensor = cutlass_torch.cute_tensor_like( + cute_f32_torch_tensor_cpu, + dtype, + is_dynamic_layout=True, + assumed_align=16, + ) + + # Convert f32 cute tensor to dtype cute tensor + cute_tensor = cutlass_torch.convert_cute_tensor( + cute_f32_torch_tensor, + cute_tensor, + dtype, + is_dynamic_layout=True, + ) + return ref_f32_torch_tensor_cpu, cute_tensor, cute_torch_tensor + + +def create_scale_factor_tensor_unswizzled( + L, mn, k, sf_vec_size, dtype, mark_dynamic_leading_dim=None, pad_sf_k=False +): + def ceil_div(a, b): + return (a + b - 1) // b + + # The M-direction stride is sf_k, and the 4-byte cp.async that walks it needs + # that to be a multiple of 4. For SFA, k is C and the scale-factor atom gate + # already makes it one, so nothing is padded. For SFD, k is Kout, which no + # gate constrains, so its caller asks for the tail. + sf_k = ceil_div(k, sf_vec_size) + alloc_sf_k = ceil_div(sf_k, 4) * 4 if pad_sf_k else sf_k + + # Use pure PyTorch to create the tensor, bypassing cute_tensor_like which + # requires the leading mode to be divisible by 4 for 8-bit types. + sf_raw = torch.randint(0, 3, (L, mn, alloc_sf_k), dtype=torch.uint8).permute( + 1, 2, 0 + ) + if alloc_sf_k > sf_k: + sf_raw[:, sf_k:, :] = 0 + + sf_torch = sf_raw.to(dtype=cutlass_torch.dtype(dtype)).cuda() + sf_tensor = from_dlpack(sf_torch, assumed_align=16) + # When the kernel indexes this tensor directly (SFA: the device reads its + # gmem through mSFA_mkl's own strides, no host rebuild), the sf_k extent + # is C/sf_vec_size and its row stride must stay runtime so one cubin + # serves every C. Without this the stride is frozen to the compile C and + # a different C reads every row at the wrong offset. + if mark_dynamic_leading_dim is not None: + sf_tensor = sf_tensor.mark_layout_dynamic(leading_dim=mark_dynamic_leading_dim) + + # Build the f32 reference for verification, over sf_k factors -- the + # allocation carries more only when the caller asked for the tail. + # Decode the scale factors through the storage dtype so the + # reference matches the values the kernel reads: E8M0 (MXFP8) has no + # exact zero (0 stores as the smallest exponent, ~5.9e-39) and only + # represents powers of two, so the raw uint8 -> f32 cast would disagree + # with the kernel on those blocks. E4M3 decodes the small integers + # losslessly, leaving the NVFP4 path unchanged. + sf_ref = ( + sf_raw[:, :sf_k, :] + .to(dtype=cutlass_torch.dtype(dtype)) + .float() + .permute(2, 0, 1) + .unsqueeze(-1) + .expand(L, mn, sf_k, sf_vec_size) + .reshape(L, mn, sf_k * sf_vec_size) + .permute(1, 2, 0) + ) + sf_ref = sf_ref[:, :k, :] + + return sf_ref, sf_tensor, sf_torch + + +def run( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + stride_dhw: Tuple[int, int, int] = (1, 1, 1), + upper_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + lower_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + dil_dhw: Tuple[int, int, int] = (1, 1, 1), + ab_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN, + d_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN, + c_dtype: Optional[Type[cutlass.Numeric]] = None, + acc_dtype: Type[cutlass.Numeric] = cutlass.Float32, + sf_dtype: Type[cutlass.Numeric] = cutlass.Float8E4M3FN, + sf_vec_size: int = 16, + cta_tile_k: Optional[int] = None, + mma_tiler_mn: Tuple[int, int] = (128, 128), + preferred_cluster_shape_mn: Tuple[int, int] = (2, 1), + fallback_cluster_shape_mn: Tuple[int, int] = (1, 1), + swizzle_size: int = 1, + raster_along: Literal["m", "n"] = "m", + use_2cta_instrs: bool = False, + tolerance: float = 1e-02, + warmup_iterations: int = 0, + iterations: int = 1, + use_cold_l2: bool = False, + skip_ref_check: bool = False, + use_bias: bool = False, + beta: float = 0.0, + activation: str = "identity", + **kwargs, +): + """Run 3D convolution and compare against PyTorch reference. + + The filter tensor uses native KTRSC layout (K, T, R, S, C) with C contiguous. + + :param ncdhw: Input tensor shape (N, C, D, H, W) + :param ktrs: Filter tensor shape components (K, T, R, S) + :param stride_dhw: Stride (Sd, Sh, Sw) + :param upper_pad_dhw: Upper padding (PadD, PadH, PadW) + :param lower_pad_dhw: Lower padding (PadD, PadH, PadW) + :param dil_dhw: Dilation (DilD, DilH, DilW) + :param ab_dtype: Data type for A/B input tensors + :param d_dtype: Data type for output tensor D + :param acc_dtype: Accumulator data type + :param mma_tiler_mn: MMA tiler shape + :param preferred_cluster_shape_mn: Preferred cluster shape (M, N) for CLC dynamic scheduling + :param fallback_cluster_shape_mn: Fallback cluster shape (M, N) for CLC dynamic scheduling + :param swizzle_size: Swizzling size in the unit of cluster. 1 means no swizzle + :param raster_along: Rasterization order of clusters + :param use_2cta_instrs: Whether to use 2CTA instructions + :param tolerance: Tolerance for result comparison + :param warmup_iterations: Number of warmup iterations + :param iterations: Number of benchmark iterations + :param use_cold_l2: Whether to flush L2 cache between iterations + :param skip_ref_check: Whether to skip reference checking + """ + if not torch.cuda.is_available(): + raise RuntimeError("GPU is required to run this example!") + + N, C, D, H, W = ncdhw + K, T, R, S = ktrs + + # Residual (C) shares the output's shape and im2col TMA descriptor, so its + # dtype defaults to the output dtype when the caller does not set one. + if c_dtype is None: + c_dtype = d_dtype + + # cta_tile_k is the compile-time K tile. When unset, take the largest legal tile + # that fits in the channel count. A C the tile does not divide is fine -- the + # trailing tile is partial. A caller may pass one explicitly to reuse a cubin + # across several C, or to trade tile width against channel padding. + # + # MXFP8 defaults to tile_k=128, the width of its SF atom, so a K tile stages + # exactly one atom. tile_k=64 is legal too -- two K tiles then share one staged + # atom -- and halves the channel padding on a C that 128 overshoots. + if cta_tile_k is None: + if ab_dtype is cutlass.Float8E4M3FN: + cta_tile_k = 128 + else: + cta_tile_k = max(t for t in (64, 128, 256) if t <= max(64, min(256, C))) + + Z, P, Q = compute_zpq( + (D, H, W), + (T, R, S), + stride_dhw, + upper_pad_dhw, + lower_pad_dhw, + dil_dhw, + ) + + print("Running Blackwell 3D Convolution test with:") + print(f" Input shape (N, C, D, H, W): {ncdhw}") + print(f" Filter shape (K, C, T, R, S): ({K}, {C}, {T}, {R}, {S})") + print(f" Output shape (N, K, Z, P, Q): ({N}, {K}, {Z}, {P}, {Q})") + print(f" Stride (Sd, Sh, Sw): {stride_dhw}") + print(f" Upper padding (PadD, PadH, PadW): {upper_pad_dhw}") + print(f" Lower padding (PadD, PadH, PadW): {lower_pad_dhw}") + print(f" Dilation (DilD, DilH, DilW): {dil_dhw}") + print(f" A/B data type: {ab_dtype}") + print(f" D data type: {d_dtype}") + print(f" Accumulator type: {acc_dtype}") + print(f" sf data type: {sf_dtype}") + print(f" sf vec size: {sf_vec_size}") + print(f" MMA tiler (M, N): {mma_tiler_mn}") + print(f" Preferred cluster shape (M, N): {preferred_cluster_shape_mn}") + print(f" Fallback cluster shape (M, N): {fallback_cluster_shape_mn}") + print(f" Swizzle size: {swizzle_size}") + print(f" Raster along: {raster_along}") + print(f" Use 2CTA instructions: {use_2cta_instrs}") + print() + + # Create input and filter tensors + input_tensor, filter_tensor, output_tensor = prepare_tensors( + ncdhw, ktrs, (Z, P, Q), ab_dtype + ) + + # Prepare cute tensors + input_, input_storage = create_cute_tensor(input_tensor, ab_dtype, leading_dim=4) + filter_, filter_storage = create_cute_tensor(filter_tensor, ab_dtype, leading_dim=4) + output_, output_storage = create_cute_tensor(output_tensor, d_dtype, leading_dim=4) + + sfa_ref, sfa_, sfa_storage = create_scale_factor_tensor_unswizzled( + 1, + N * D * H * W, + C, + sf_vec_size, + sf_dtype, + mark_dynamic_leading_dim=1, + ) + # Allocate SFB at the per-position span: the channels past C pair with the + # zero-filled B the TMA produces there, so what they scale is zero. + sfb_c_span = sfb_per_position_channels(C, cta_tile_k, sf_vec_size) + sfb_ref, sfb_, sfb_storage = create_scale_factor_tensor_swizzled( + 1, + K, + sfb_c_span * T * R * S, + sf_vec_size, + sf_dtype, + ) + # SFD: emitted for a narrow (block-scaled) output; FP16 and BF16 carry none. It + # takes the input's scale format outright, so NVFP4 stays E4M3 over 16 channels + # and MXFP8 E8M0 over 32 on both sides. + gen_sfd = d_dtype in (cutlass.Float4E2M1FN, cutlass.Float8E4M3FN) + sfd_vec_size = sf_vec_size + sfd_dtype = sf_dtype + if gen_sfd: + _, sfd_, sfd_storage = create_scale_factor_tensor_unswizzled( + 1, + N * Z * P * Q, + K, + sfd_vec_size, + sfd_dtype, + pad_sf_k=True, + ) + # Helper randomizes by default -- we want to read what the kernel writes. + sfd_storage.zero_() + else: + sfd_, sfd_storage = None, None + + # Per-output-channel bias (length K), dtype matches the output. Added to the + # accumulator in FP32 as D = alpha*acc + bias, broadcast across every spatial + # output position. + if use_bias: + bias_storage = torch.randn(K, dtype=torch.float32, device="cuda").to( + cutlass_torch.dtype(d_dtype) + ) + bias_ = from_dlpack(bias_storage, assumed_align=16) + else: + bias_storage, bias_ = None, None + + # Per-element residual (C), same (N, Z, P, Q, K) shape as the output D. + # Added to the accumulator in FP32 as D = alpha*acc + bias + beta*C. + # Built as a dynamic-layout cute tensor (leading_dim=4, K contiguous) so it + # is described by the same im2col TMA descriptor as the output store. + if beta != 0.0: + c_f32_src = torch.randn((N, Z, P, Q, K), dtype=torch.float32, device="cuda") + c_, c_storage = create_cute_tensor(c_f32_src, c_dtype, leading_dim=4) + else: + c_f32_src, c_storage, c_ = None, None, None + + # Per-tensor FP32 scaling factors, passed as runtime scalars: + # alpha: applied in accumulator precision before quantize + # norm_const: folded into SFD scale-up (only used when gen_sfd) + alpha_ = cutlass.Float32(1.0) + norm_const_ = cutlass.Float32(1.0) + + # Resolve the compile-time epilogue activation. The device op is folded into + # the cubin as a Constexpr; the reference op mirrors it on the host. + if activation not in EPILOGUE_ACTIVATIONS: + raise ValueError( + f"Unsupported activation {activation!r}; " + f"choose from {sorted(EPILOGUE_ACTIVATIONS)}" + ) + epilogue_op = EPILOGUE_ACTIVATIONS[activation]["device"] + + # Compile convolution kernel + print("Compiling kernel with cute.compile ...") + compiled_fn = compile_conv( + ncdhw, + ktrs, + input_, + filter_, + output_, + acc_dtype, + sfa_, + sfb_, + sfd_, + alpha_, + norm_const_, + bias_, + c_, + beta, + sf_vec_size, + cta_tile_k, + mma_tiler=mma_tiler_mn, + preferred_cluster_shape_mn=preferred_cluster_shape_mn, + fallback_cluster_shape_mn=fallback_cluster_shape_mn, + swizzle_size=swizzle_size, + raster_along=raster_along, + use_2cta_instrs=use_2cta_instrs, + upper_padding_dhw=upper_pad_dhw, + lower_padding_dhw=lower_pad_dhw, + stride_dhw=stride_dhw, + dilation_dhw=dil_dhw, + epilogue_op=epilogue_op, + ) + + # Get current CUDA stream + torch_stream = torch.cuda.Stream() + current_stream = cuda.CUstream(torch_stream.cuda_stream) + + # The compiled entry expects the 12 pad/stride/dilation runtime scalars + # (upper d/h/w, lower d/h/w, stride d/h/w, dil d/h/w) matching the boxed + # cutlass.Int32 params baked into cute.compile; pass them as the same host + # config values so one cubin serves any config. + print("Running Blackwell 3D convolution...") + # The inputs initialize on other streams; drain them so the kernel's + # stream never reads a tensor mid-initialization. + torch.cuda.synchronize() + compiled_fn( + input_, + filter_, + output_, + sfa_, + sfb_, + alpha_, + sfd_, + norm_const_, + bias_, + c_, + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + current_stream, + ) + torch_stream.synchronize() + + if not skip_ref_check: + print("Verifying results with block-scaled reference...") + # Block-scaled reference using F.conv3d with pre-scaled inputs + # This handles padding/stride/dilation correctly + + # Pre-scale input by SFA: sfa_ref shape (N*D*H*W, C, 1) + sfa_expanded = sfa_ref.squeeze(-1).reshape(N, D, H, W, C).cuda() + scaled_input = input_tensor.cuda().float() * sfa_expanded + + # Pre-scale filter by SFB: sfb_ref shape (K, C*T*R*S, 1) + # Drop the per-position padding channels: they scale zero-filled B and have + # no counterpart in the reference convolution. + sfb_expanded = ( + sfb_ref.squeeze(-1).reshape(K, T, R, S, sfb_c_span)[..., :C].cuda() + ) + scaled_filter = filter_tensor.cuda().float() * sfb_expanded + + # F.conv3d expects (N, C, D, H, W) and (K, C, T, R, S) + scaled_input_ncdhw = scaled_input.permute(0, 4, 1, 2, 3).contiguous() + scaled_filter_kctrs = scaled_filter.permute(0, 4, 1, 2, 3).contiguous() + + if upper_pad_dhw == lower_pad_dhw: + ref_nkzpq = F.conv3d( + scaled_input_ncdhw, + scaled_filter_kctrs, + padding=upper_pad_dhw, + stride=stride_dhw, + dilation=dil_dhw, + ) + else: + # Asymmetric padding: F.pad then conv3d without padding. + # Convention: lower_pad = leading (left), upper_pad = trailing (right). + # This must match the im2col TMA descriptor which subtracts lower_pad + # when mapping output coords to input coords (d_in = z*stride - lower_pad + t*dil). + padded = F.pad( + scaled_input_ncdhw, + ( + lower_pad_dhw[2], + upper_pad_dhw[2], + lower_pad_dhw[1], + upper_pad_dhw[1], + lower_pad_dhw[0], + upper_pad_dhw[0], + ), + ) + ref_nkzpq = F.conv3d( + padded, + scaled_filter_kctrs, + stride=stride_dhw, + dilation=dil_dhw, + ) + + # Convert to NZPQK layout (matching kernel output) + ref = ref_nkzpq.permute(0, 2, 3, 4, 1).contiguous() # (N, Z, P, Q, K) + # Add per-output-channel bias (D = alpha*acc + bias) in the FP32 domain, + # broadcast across N/Z/P/Q. alpha is 1.0 here so acc == conv result. + if use_bias: + ref = ref + bias_storage.float().view(1, 1, 1, 1, K) + # Add the per-element residual (D = alpha*acc + bias + beta*residual) in + # the FP32 domain. c_storage holds the exact bf16 values the + # kernel loads, so float() reproduces them bit-for-bit. + if beta != 0.0: + ref = ref + beta * c_storage.float().reshape(N, Z, P, Q, K).cuda() + # Apply the epilogue activation on the full linear combination + # (D = activation(alpha*acc + bias + beta*residual)), matching the + # device op folded into the kernel. + ref = EPILOGUE_ACTIVATIONS[activation]["ref"](ref) + # Snapshot the un-quantized FP32 ref BEFORE the in-place quantize round-trip + # below mutates `ref` (shares GPU storage with ref_device). + ref_unquant_cpu = ref.detach().cpu().clone() + + # Convert kernel FP4 output to f32 for comparison + d_ref_device = torch.empty((N, Z, P, Q, K), dtype=torch.float32, device="cuda") + cute.testing.convert( + output_, + from_dlpack(d_ref_device, assumed_align=16).mark_layout_dynamic( + leading_dim=4 + ), + ) + d_ref_result = d_ref_device.cpu() + + # Quantize reference: f32 -> d_dtype -> f32 (mutates ref in-place via shared + # storage). Scratch storage dtype follows the kernel-convert rule: sub-byte + # FP4 and <=8-bit floats pack into a uint8 byte buffer, while wider types + # (f16/bf16/f32) use their native torch dtype so the buffer is not + # under-allocated (f16 needs 2 bytes/elem, not 1). + ref_quant_byte_storage = (d_dtype.is_float and d_dtype.width <= 8) or ( + d_dtype.is_integer and d_dtype.width == 4 + ) + ref_quant_storage_dtype = ( + torch.uint8 if ref_quant_byte_storage else cutlass_torch.dtype(d_dtype) + ) + ref_f4_ = torch.empty( + (N, Z, P, Q, K), dtype=ref_quant_storage_dtype, device="cuda" + ) + ref_f4 = from_dlpack(ref_f4_, assumed_align=16).mark_layout_dynamic( + leading_dim=4 + ) + ref_f4.element_type = d_dtype + ref_device = ref.contiguous().cuda() + ref_tensor = from_dlpack(ref_device, assumed_align=16).mark_layout_dynamic( + leading_dim=4 + ) + cute.testing.convert(ref_tensor, ref_f4) + cute.testing.convert(ref_f4, ref_tensor) + ref_quantized = ref_device.cpu() + + if gen_sfd: + # SFD path: the kernel rescales acc by norm_const / SFD before the + # FP4 cast, so the raw kernel D is only meaningful together with SFD. + # Recompute the reference SFD and the SFD-rescaled reference output + # fully on the host -- per-vector amax, scale factor, and rescale are + # all derived from the un-quantized reference, independent of the + # kernel's own SFD -- then compare SFD and the quantized output + # elementwise. The reference block scale is read back after its E4M3 + # cast so the rescale is self-consistent, matching the kernel. + sf_k = (K + sfd_vec_size - 1) // sfd_vec_size + # norm_const_ is a device Float32; mirror it as a host float for the + # torch-side replay (host path always uses the 1.0 default). + norm_const_host = 1.0 + fp32_max = torch.finfo(torch.float32).max + + # 1. Reference SFD and SFD-rescaled reference output, host-side. + # per-vector amax over sfd_vec_size contiguous K elements + # -> pvscale = amax * norm_const / M_D + # -> cast to sfd_dtype and read back (sfd_ref) + # -> rescale ref by norm_const / sfd_ref. + # M_D is the output dtype's full-scale max (FP4=6.0, E4M3=448.0); + # sfd_dtype is the input sf_dtype (E4M3 for NVFP4, E8M0 for MXFP8). + m_d = 6.0 if d_dtype is cutlass.Float4E2M1FN else 448.0 + sfd_torch_dtype = cutlass_torch.dtype(sfd_dtype) + k_pad = sf_k * sfd_vec_size + ref_pad = torch.zeros((N, Z, P, Q, k_pad), dtype=torch.float32) + ref_pad[..., :K] = ref_unquant_cpu + amax = ( + ref_pad.reshape(N, Z, P, Q, sf_k, sfd_vec_size).abs().amax(dim=5) + ) # (N, Z, P, Q, sf_k) + pvscale = amax * norm_const_host / m_d + if sfd_dtype is cutlass.Float8E8M0FNU: + # E8M0 SFD (MXFP8) rounds toward +inf to the next power of two, + # matching the device cvt.rp.ue8m0; a plain round-to-nearest cast + # would disagree by up to one exponent step. exp2/log2 run on CPU + # (the CUDA path goes through nvrtc jiterator). E8M0 has no zero; + # its smallest value is 2^-127, which cvt.rp of 0.0 also yields. + sfd_ref = torch.exp2( + torch.ceil(torch.log2(pvscale)).clamp(-127.0, 127.0) + ) + else: + # E4M3 SFD (NVFP4) has a mantissa; round-to-nearest is exact enough. + sfd_ref = pvscale.to(sfd_torch_dtype).to(torch.float32) + acc_scale = (norm_const_host / sfd_ref.clamp(min=1e-30)).clamp(max=fp32_max) + d_scaled = ( + ref_unquant_cpu + * acc_scale.unsqueeze(-1) + .expand(N, Z, P, Q, sf_k, sfd_vec_size) + .reshape(N, Z, P, Q, k_pad)[..., :K] + ) + + # 2. Quantize the rescaled reference through d_dtype (f32 -> FP4 -> + # f32) via the same convert path the kernel output went through. + d_ref_narrow_ = torch.empty( + (N, Z, P, Q, K), dtype=ref_quant_storage_dtype, device="cuda" + ) + d_ref_narrow = from_dlpack( + d_ref_narrow_, assumed_align=16 + ).mark_layout_dynamic(leading_dim=4) + d_ref_narrow.element_type = d_dtype + d_scaled_dev = d_scaled.contiguous().cuda() + d_scaled_tensor = from_dlpack( + d_scaled_dev, assumed_align=16 + ).mark_layout_dynamic(leading_dim=4) + cute.testing.convert(d_scaled_tensor, d_ref_narrow) + cute.testing.convert(d_ref_narrow, d_scaled_tensor) + d_ref_quant = d_scaled_dev.cpu() + + # 3. Decode the kernel-written SFD (plain (NZPQ, sf_k_padded, 1) + # layout) to f32 in the same (N, Z, P, Q, sf_k) shape as sfd_ref, + # reinterpreting the raw bytes through the actual SFD dtype. + sfd_cpu = sfd_storage.cpu() + sfd_bytes_view = sfd_cpu.view(torch.uint8) + sfd_kernel = ( + sfd_bytes_view[:, :sf_k, :] + .view(sfd_torch_dtype) + .float() + .squeeze(-1) + .reshape(N, Z, P, Q, sf_k) + ) + n_nonzero = (sfd_bytes_view[:, :sf_k, :] != 0).sum().item() + assert n_nonzero > 0, "SFD is all zero -- kernel did not write SFD" + + # Compare SFD first: the output rescale depends on it, so an SFD + # mismatch is the more fundamental signal. + torch.testing.assert_close(sfd_kernel, sfd_ref, atol=1e-01, rtol=1e-01) + torch.testing.assert_close( + d_ref_result, d_ref_quant, atol=1e-01, rtol=1e-01 + ) + print("SFD dequant refcheck passed.") + else: + torch.testing.assert_close( + d_ref_result, + ref_quantized, + atol=tolerance, + rtol=1e-02, + ) + print("Results match within tolerance!") + + # Benchmark if requested + if iterations > 0: + print( + f"\nBenchmarking with {warmup_iterations} warmup and {iterations} iterations..." + ) + + def generate_tensors(): + input_tensor, filter_tensor, output_tensor = prepare_tensors( + ncdhw, ktrs, (Z, P, Q), ab_dtype + ) + input_, input_storage = create_cute_tensor( + input_tensor, ab_dtype, leading_dim=4 + ) + filter_, filter_storage = create_cute_tensor( + filter_tensor, ab_dtype, leading_dim=4 + ) + output_, output_storage = create_cute_tensor( + output_tensor, d_dtype, leading_dim=4 + ) + # Arg order must match the compiled entry exactly (epilogue_op and + # beta are Constexpr, folded in at cute.compile time, so they are + # NOT runtime args): a, b, d, sfa, sfb, alpha, sfd, norm_const, + # bias, residual, then the 12 pad/stride/dil runtime Int32, then + # stream last. Only the A/B/D tensors rotate per cold-L2 workspace; + # the SF tensors, alpha/norm_const/bias/residual, and the geometry + # scalars are reused. + # The workspace initializes on other streams; drain them so the + # benchmark stream never reads a tensor mid-initialization. + torch.cuda.synchronize() + return testing.JitArguments( + input_, + filter_, + output_, + sfa_, + sfb_, + alpha_, + sfd_, + norm_const_, + bias_, + c_, + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + current_stream, + ) + + workspace_count = 1 + if use_cold_l2: + one_workspace_bytes = ( + input_storage.numel() * input_storage.element_size() + + filter_storage.numel() * filter_storage.element_size() + + output_storage.numel() * output_storage.element_size() + + sfa_storage.numel() * sfa_storage.element_size() + + sfb_storage.numel() * sfb_storage.element_size() + ) + workspace_count = testing.get_workspace_count( + one_workspace_bytes, warmup_iterations, iterations + ) + + # exec_time is in microseconds + exec_time = testing.benchmark( + compiled_fn, + workspace_generator=generate_tensors, + workspace_count=workspace_count, + stream=current_stream, + warmup_iterations=warmup_iterations, + iterations=iterations, + use_cuda_graphs=True, + ) + runtime_s = exec_time / 1.0e6 + fmas = (N * Z * P * Q) * K * (C * T * R * S) + flop = 2 * fmas + gflop = flop / 1.0e9 + gflops = gflop / runtime_s + + print("Average Runtime : ", exec_time / 1000, "ms") + print("GFLOPS : ", gflops) + + return exec_time + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="Blackwell 3D convolution") + + # Convolution parameters + parser.add_argument( + "--ncdhw", + type=_parse_comma_separated_ints, + default=(1, 128, 32, 32, 32), + help="Input tensor shape (N,C,D,H,W)", + ) + parser.add_argument( + "--ktrs", + type=_parse_comma_separated_ints, + default=(256, 3, 3, 3), + help="Filter tensor shape components (K,T,R,S)", + ) + parser.add_argument( + "--stride_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Stride (Sd,Sh,Sw)", + ) + parser.add_argument( + "--upper_pad_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Upper padding (PadD,PadH,PadW)", + ) + parser.add_argument( + "--lower_pad_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Lower padding (PadD,PadH,PadW)", + ) + parser.add_argument( + "--dil_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Dilation (DilD,DilH,DilW)", + ) + + # Data type parameters + parser.add_argument( + "--ab_dtype", + type=cutlass.dtype, + choices=[ + cutlass.Float4E2M1FN, + cutlass.Float8E4M3FN, + ], + default=cutlass.Float4E2M1FN, + help="Data type for A/B input tensors", + ) + parser.add_argument( + "--d_dtype", + type=cutlass.dtype, + choices=[ + cutlass.Float4E2M1FN, + cutlass.Float8E5M2, + cutlass.Float8E4M3FN, + cutlass.Float16, + cutlass.BFloat16, + cutlass.Float32, + ], + default=cutlass.Float4E2M1FN, + help="Output D data type. SFD is generated only when width <= 8.", + ) + parser.add_argument( + "--c_dtype", + type=cutlass.dtype, + default=None, + help="Residual (C) input data type. Defaults to --d_dtype; must equal " + "it since the residual shares the output's im2col TMA descriptor.", + ) + parser.add_argument( + "--acc_dtype", + type=cutlass.dtype, + choices=[cutlass.Float32, cutlass.Float16, cutlass.Int32], + default=cutlass.Float32, + help="Accumulator data type", + ) + parser.add_argument( + "--sf_dtype", + type=cutlass.dtype, + choices=[ + cutlass.Float8E4M3FN, + cutlass.Float8E8M0FNU, + ], + default=cutlass.Float8E4M3FN, + help="Data type for A/B/D scaling factor tensors (NVFP4 default: E4M3)", + ) + parser.add_argument("--sf_vec_size", type=int, default=16) + parser.add_argument( + "--cta_tile_k", + type=int, + default=None, + help="Compile-time K tile (64/128/192/256). Defaults to min(256, C).", + ) + + # Kernel parameters + parser.add_argument( + "--mma_tiler_mn", + type=_parse_comma_separated_ints, + default=(128, 128), + help="MMA tiler shape (M,N)", + ) + parser.add_argument( + "--preferred_cluster_shape_mn", + type=_parse_comma_separated_ints, + default=(2, 1), + help="Preferred cluster shape (M,N) for CLC dynamic scheduling", + ) + parser.add_argument( + "--fallback_cluster_shape_mn", + type=_parse_comma_separated_ints, + default=(2, 1), + help="Fallback cluster shape (M,N) for CLC dynamic scheduling", + ) + parser.add_argument( + "--use_2cta_instrs", + action="store_true", + help="Enable 2CTA MMA instructions", + ) + parser.add_argument( + "--swizzle_size", + type=int, + default=1, + help="Swizzling size in the unit of cluster for improving L2 cache hit rate", + ) + parser.add_argument( + "--raster_order", + type=str, + choices=["m", "n"], + default="m", + help="Rasterization order of clusters", + ) + # Testing parameters + parser.add_argument( + "--tolerance", + type=float, + default=1e-02, + help="Tolerance for result comparison", + ) + parser.add_argument( + "--warmup_iterations", + type=int, + default=0, + help="Number of warmup iterations", + ) + parser.add_argument( + "--iterations", + type=int, + default=0, + help="Number of benchmark iterations", + ) + parser.add_argument( + "--skip_ref_check", + action="store_true", + help="Skip reference checking", + ) + parser.add_argument( + "--use_cold_l2", + action="store_true", + default=False, + help="Use circular buffer tensor sets to ensure L2 cold cache", + ) + parser.add_argument( + "--use_bias", + action="store_true", + default=False, + help="Add a per-output-channel bias (D = alpha*acc + bias)", + ) + parser.add_argument( + "--beta", + type=float, + default=0.0, + help="Residual scaling (D = alpha*acc + bias + beta*residual); " + "beta != 0 enables the residual path, beta == 0 disables it", + ) + parser.add_argument( + "--activation", + type=str, + default="identity", + choices=sorted(EPILOGUE_ACTIVATIONS), + help="Compile-time epilogue activation applied as " + "D = activation(alpha*acc + bias + beta*residual). One activation per " + "cubin (folded in as a Constexpr). Non-identity is unsupported for FP4 " + "output.", + ) + args = parser.parse_args() + + run( + args.ncdhw, + args.ktrs, + args.stride_dhw, + args.upper_pad_dhw, + args.lower_pad_dhw, + args.dil_dhw, + args.ab_dtype, + args.d_dtype, + args.c_dtype, + args.acc_dtype, + args.sf_dtype, + args.sf_vec_size, + args.cta_tile_k, + args.mma_tiler_mn, + args.preferred_cluster_shape_mn, + args.fallback_cluster_shape_mn, + args.swizzle_size, + args.raster_order, + args.use_2cta_instrs, + args.tolerance, + args.warmup_iterations, + args.iterations, + args.use_cold_l2, + args.skip_ref_check, + args.use_bias, + args.beta, + args.activation, + ) + print("PASS") diff --git a/examples/python/CuTeDSL/cute/blackwell/kernel/conv/dense_implicit_gemm_fprop.py b/examples/python/CuTeDSL/cute/blackwell/kernel/conv/dense_implicit_gemm_fprop.py new file mode 100644 index 0000000000..2f310aeef3 --- /dev/null +++ b/examples/python/CuTeDSL/cute/blackwell/kernel/conv/dense_implicit_gemm_fprop.py @@ -0,0 +1,2721 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause + +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: + +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. + +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. + +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. + +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import argparse +from typing import Optional, Tuple, Type, Union, Literal +from functools import lru_cache +import cuda.bindings.driver as cuda +import sys +import os + +import torch +import torch.nn.functional as F + +import cutlass +import cutlass.cute as cute +from cutlass import testing +from cutlass.cute.runtime import from_dlpack +from cutlass.torch import dtype as torch_dtype +import cutlass.utils as utils +from cutlass.utils import is_fp8_dtype +from cutlass.cute.nvgpu import cpasync, tcgen05 +import cutlass.utils.blackwell_helpers as sm100_utils +import cutlass.pipeline as pipeline +from cutlass.pipeline import pipeline_init_arrive, pipeline_init_wait +from pathlib import Path + +if __name__ == "__main__": + # `helpers` sits at the examples/CuTeDSL root; running this file as a + # script only puts its own directory on sys.path. + cutedsl_dir = str(Path(__file__).resolve().parents[4]) + if cutedsl_dir not in sys.path: + sys.path.insert(0, cutedsl_dir) + +from cutlass.utils import ( + ClcDynamicPersistentTileScheduler, + ClcDynamicPersistentTileSchedulerParams, +) + +if __name__ == "__main__": + current_dir = os.path.dirname(os.path.abspath(__file__)) + sys.path.insert(0, os.path.join(current_dir, "../../../")) + +from blackwell.kernel.dense_gemm.dense_gemm_persistent_dynamic import ( + PersistentDenseGemmKernel, + _compute_stages, +) + +""" +A high-performance 3D implicit-GEMM based fprop convolution example for the NVIDIA Blackwell SM100 architecture using CUTE DSL. +- Input tensor A is NxDxHxWxC, must be C major. +- Filter tensor B is KxTxRxSxC, must be C major. +- Output tensor D is NxZxPxQxK, must be K major. + +This kernel supports the following features: + - Utilizes Tensor Memory Access (TMA) for efficient memory operations and for the im2col transformation of the input tensor A. + - Utilizes Blackwell's tcgen05 MMA for matrix multiply-accumulate (MMA) operations (including 2cta mma instructions) + - Implements TMA multicast with cluster to reduce L2 memory traffic + - Utilizes a CLC dynamic persistent dense GEMM kernel for the implicit GEMM. + +This implicit-GEMM based convolution works by converting the convolution into a GEMM problem with the following mapping: +- GEMM M dimension maps to NxZxPxQ +- GEMM N dimension maps to K +- GEMM K dimension maps to TxRxSxC +During the load of input tensor to SMEM, the TMA operation performs the im2col transformation on the input tensor A. +This transforms the A matrix into the required shape for the GEMM operation (NxZxPxQ by TxRxSxC), and may involve replication of the input elements. +Filter tensor can be loaded to SMEM without any transformation. +The output tensor D is then stored to GMEM via the im2col TMA store (no transformation necessary). + +To run this example: + +.. code-block:: bash + + python examples/CuTeDSL/cute/blackwell/kernel/conv/dense_implicit_gemm_fprop.py \ + --ncdhw 1,128,32,32,32 --ktrs 256,3,3,3 \ + --ab_dtype Float16 --d_dtype Float16 --acc_dtype Float32 \ + --use_2cta_instrs --mma_tiler_mn 256,128 \ + --preferred_cluster_shape_mn 2,1 --fallback_cluster_shape_mn 2,1 \ + --upper_pad_dhw 1,1,1 --lower_pad_dhw 1,1,1 \ + --stride_dhw 1,1,1 --dil_dhw 1,1,1 + +To collect performance with NCU profiler: + +.. code-block:: bash + + ncu python examples/CuTeDSL/cute/blackwell/kernel/conv/dense_implicit_gemm_fprop.py \ + --ncdhw 1,128,32,32,32 --ktrs 256,3,3,3 \ + --ab_dtype Float16 --d_dtype Float16 --acc_dtype Float32 \ + --use_2cta_instrs --mma_tiler_mn 256,128 \ + --preferred_cluster_shape_mn 2,1 --fallback_cluster_shape_mn 2,1 \ + --upper_pad_dhw 1,1,1 --lower_pad_dhw 1,1,1 \ + --stride_dhw 1,1,1 --dil_dhw 1,1,1 \ + --warmup_iterations 1 --iterations 10 --skip_ref_check + +Constraints: +* Supported input data types: fp16, bf16, tf32, int8, uint8, fp8 (e4m3fn, e5m2), + see detailed valid dtype combinations in below + Sm100PersistentDenseImplicitGemmFpropKernel class documentation +* A/B tensor must have the same data type +* Mma tiler M must be 64/128 (use_2cta_instrs=False) or 128/256 (use_2cta_instrs=True) +* Mma tiler N must be 32-256, step 32 +* Cluster shape M/N must be positive and power of 2, total cluster size <= 16 +* Cluster shape M must be multiple of 2 if use_2cta_instrs=True +* Preferred cluster shape M/N must be positive and power of 2, total cluster size <= 16 +* Preferred cluster shape M must be multiple of 2 if use_2cta_instrs=True +* Preferred cluster shape M/N must be multiple of fallback cluster shape M/N +* The contiguous dimension of A/B/D tensors must be at least 16 bytes aligned, + i.e, number of elements is a multiple of 4, 8, and 16 for TFloat32, + Float16/BFloat16, and Int8/Uint8/Float8, respectively. +""" + + +def _check_tensor_alignment( + c: int, + k: int, + ab_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], +): + """Check if the tensor alignment is valid for convolution. + + :param c: The number of input channels + :type c: int + :param k: The number of output channels + :type k: int + :param ab_dtype: The data type of the A and B operands + :type ab_dtype: Type[cutlass.Numeric] + :param d_dtype: The data type of the output tensor + :type d_dtype: Type[cutlass.Numeric] + """ + + def check_contiguous_16B_alignment(dtype, num_major_elements): + num_contiguous_elements = 16 * 8 // dtype.width + return num_major_elements % num_contiguous_elements == 0 + + if not check_contiguous_16B_alignment( + d_dtype, k + ) or not check_contiguous_16B_alignment(ab_dtype, c): + raise testing.CantImplementError( + f"Invalid tensor alignment: C = {c}, K = {k}, ab_dtype = {ab_dtype}, d_dtype = {d_dtype}" + ) + + +def _check_im2col_descriptor_limits( + filter_trs: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + dilation_dhw: Tuple[int, int, int], +): + """Check that the convolution geometry fits the im2col tensor map's fields. + + Three fields of the 5D im2col tensor map are narrower than the convolution + parameters feeding them, and a 3D convolution always builds a 5D descriptor. The + widths come from the encodings: a 5-bit signed corner per spatial dimension + ([-16, 15]), an unsigned 5-bit coordinate offset per dimension ([0, 31]), and a + 3-bit traversal stride holding the stride minus one ([1, 8]). A corner is not the + padding itself -- the leading one is -lower_padding, the trailing one + upper_padding - (filter - 1) * dilation -- so padding and dilation only bind in + combination. Overflowing one truncates it and moves the pixel box, surfacing as an + illegal instruction when the shifted box leaves mapped memory and as a wrong answer + when it does not, so computing correctly past a bound is luck, not contract. + + :param filter_trs: Filter extents (T, R, S) + :param upper_padding_dhw: Padding before the data per spatial dimension + :param lower_padding_dhw: Padding after the data per spatial dimension + :param stride_dhw: Convolution stride per spatial dimension + :param dilation_dhw: Filter dilation per spatial dimension + """ + corner_lo = -16 + corner_hi = 15 + max_filter_offset = 31 + max_element_stride = 8 + for dim, flt, dil, pad_up, pad_lo, stride in zip( + ("D", "H", "W"), + filter_trs, + dilation_dhw, + upper_padding_dhw, + lower_padding_dhw, + stride_dhw, + strict=True, + ): + leading_corner = -pad_lo + trailing_corner = pad_up - (flt - 1) * dil + if not corner_lo <= leading_corner <= corner_hi: + raise testing.CantImplementError( + f"{dim} leading im2col corner is -lower_padding = " + f"{leading_corner}, outside the [{corner_lo}, {corner_hi}] the " + f"descriptor's signed 5-bit corner encodes; lower_padding_" + f"{dim.lower()} must be at most {-corner_lo}" + ) + if not corner_lo <= trailing_corner <= corner_hi: + raise testing.CantImplementError( + f"{dim} trailing im2col corner is upper_padding - (filter - 1) " + f"* dilation = {pad_up} - ({flt} - 1) * {dil} = " + f"{trailing_corner}, outside the [{corner_lo}, {corner_hi}] the " + f"descriptor's signed 5-bit corner encodes" + ) + if stride > max_element_stride: + raise testing.CantImplementError( + f"stride_{dim.lower()}={stride} exceeds the {max_element_stride}" + f" a traversal stride encodes, holding the stride minus one in " + f"3 bits" + ) + # A filter offset is dilation * tap index, so the largest one a dimension + # reaches is (filter - 1) * dilation. Overflowing the unsigned field shifts + # the filter taps and silently changes the result, so it is bounded even when + # both corners are in range. + filter_offset = (flt - 1) * dil + if filter_offset > max_filter_offset: + raise testing.CantImplementError( + f"{dim} filter offset is (filter - 1) * dilation = ({flt} - 1) " + f"* {dil} = {filter_offset}, past the {max_filter_offset} an " + f"unsigned 5-bit im2col coordinate offset encodes" + ) + + +def _check_swizzle_size( + m: int, + n: int, + mma_tiler_mn: Tuple[int, int], + use_2cta_instrs: bool, + preferred_cluster_shape_mn: Tuple[int, int], + fallback_cluster_shape_mn: Tuple[int, int], + swizzle_size: int, + raster_along: str, +): + """Check that swizzle_size does not exceed the cluster count in the swizzled dimension. + + :param m: GEMM M dimension (N*Z*P*Q for convolution) + :param n: GEMM N dimension (K for convolution) + :param mma_tiler_mn: MMA tiler shape (M, N) + :param use_2cta_instrs: Whether 2-CTA instructions are used + :param preferred_cluster_shape_mn: Preferred cluster shape (M, N) + :param fallback_cluster_shape_mn: Fallback cluster shape (M, N) + :param swizzle_size: Swizzle size to validate + :param raster_along: Rasterization order ("m" or "n") + """ + if swizzle_size <= 1: + return + + cta_v = 2 if use_2cta_instrs else 1 + cta_tile_m = mma_tiler_mn[0] // cta_v + cta_tile_n = mma_tiler_mn[1] + m_tiles = -(-m // cta_tile_m) + n_tiles = -(-n // cta_tile_n) + + for cs in [preferred_cluster_shape_mn, fallback_cluster_shape_mn]: + if raster_along == "m": + nclusters = -(-n_tiles // cs[1]) + else: + nclusters = -(-m_tiles // cs[0]) + if nclusters < swizzle_size: + dim_name = "N" if raster_along == "m" else "M" + raise testing.CantImplementError( + f"swizzle_size ({swizzle_size}) exceeds the number of " + f"{dim_name} clusters ({nclusters}) for cluster shape " + f"{cs}. Use a smaller swizzle_size or increase the " + f"{dim_name} dimension." + ) + + +class Sm100PersistentDenseImplicitGemmFpropKernel(PersistentDenseGemmKernel): + """ + Persistent 3D convolution kernel. + The input (A) is expected to be in 5D tensor (NDHWC) format and is loaded via the im2col TMA load atom. + The filter (B) is expected to be in 5D tensor (KTRSC) format and is loaded via the TMA load atom. + The output (D) is expected to be in 5D tensor (NZPQK) format and is stored via the im2col TMA store. + The implicit GEMM runs as a megakernel that dispatches on the cluster shape the + launch actually formed, so a single launch serves both the preferred and the + fallback cluster. + + :param acc_dtype: Data type for accumulation during computation + :type acc_dtype: type[cutlass.Numeric] + :param use_2cta_instrs: Whether to use CTA group 2 for advanced thread cooperation + :type use_2cta_instrs: bool + :param mma_tiler_mn: Shape of the Matrix Multiply-Accumulate (MMA) tiler (M,N) + :type mma_tiler_mn: Tuple[int, int] + :param preferred_cluster_shape_mn: Preferred cluster dimensions (M,N) for optimal performance + :type preferred_cluster_shape_mn: Tuple[int, int] + :param fallback_cluster_shape_mn: Fallback cluster dimensions (M,N) for parallel processing + :type fallback_cluster_shape_mn: Tuple[int, int] + :param filter_trs: Filter dimensions (T, R, S) + :type filter_trs: Tuple[int, int, int] + :param upper_padding_dhw: Upper padding (PadD, PadH, PadW) + :type upper_padding_dhw: Tuple[int, int, int] + :param lower_padding_dhw: Lower padding (PadD, PadH, PadW) + :type lower_padding_dhw: Tuple[int, int, int] + :param stride_dhw: Stride (Sd, Sh, Sw) + :type stride_dhw: Tuple[int, int, int] + :param dilation_dhw: Dilation (DilD, DilH, DilW) + :type dilation_dhw: Tuple[int, int, int] + :param swizzle_size: Swizzling size in the unit of cluster. 1 means no swizzle + :type swizzle_size: int + :param raster_along: Rasterization order of clusters. Only used when swizzle_size > 1. + :type raster_along: Literal["m", "n"] + + :note: In current version, A and B tensor must be C major. D tensor must be K major. + + :note: In current version, A and B tensor must have the same data type + - i.e., Float8E4M3FN for A and Float8E5M2 for B is not supported + + :note: Supported A/B data types: + - TFloat32 + - Float16/BFloat16 + - Int8/Uint8 + - Float8E4M3FN/Float8E5M2 + + :note: Supported accumulator data types: + - Float32 (for all floating point A/B data types) + - Float16 (only for fp16 and fp8 A/B data types) + - Int32 (only for uint8/int8 A/B data types) + + :note: Supported D data types: + - Float32 (for float32 and int32 accumulator data types) + - Int32 (for float32 and int32 accumulator data types) + - Float16/BFloat16 (for fp16 and fp8 accumulator data types) + - Int8/Uint8 (for uint8/int8 accumulator data types) + - Float8E4M3FN/Float8E5M2 (for float32 accumulator data types) + + :note: Constraints: + - MMA tiler M must be 64/128 (use_2cta_instrs=False) or 128/256 (use_2cta_instrs=True) + - MMA tiler N must be 32-256, step 32 + - Cluster shape M must be multiple of 2 if use_2cta_instrs=True + - Cluster shape M/N must be positive and power of 2, total cluster size <= 16 + - Preferred cluster shape M must be multiple of 2 if use_2cta_instrs=True + - Preferred cluster shape M/N must be positive and power of 2, total cluster size <= 16 + - Preferred cluster shape M/N must be multiple of fallback cluster shape M/N + + **Example:** + + .. code-block:: python + conv = Sm100PersistentDenseImplicitGemmFpropKernel( + acc_dtype=cutlass.Float32, + use_2cta_instrs=True, + mma_tiler_mn=(128, 128), + preferred_cluster_shape_mn=(4, 2), + fallback_cluster_shape_mn=(2, 1), + ) + conv( + a, b, d, + upper_pad_d, upper_pad_h, upper_pad_w, + lower_pad_d, lower_pad_h, lower_pad_w, + stride_d, stride_h, stride_w, + dilation_d, dilation_h, dilation_w, + stream, + bias, # optional length-K bias tensor, or None + epilogue_op, + ) + """ + + def __init__( + self, + acc_dtype: Type[cutlass.Numeric], + use_2cta_instrs: bool, + mma_tiler_mn: Tuple[int, int], + preferred_cluster_shape_mn: Tuple[int, int], + fallback_cluster_shape_mn: Tuple[int, int], + swizzle_size: int = 1, + raster_along: Literal["m", "n"] = "m", + ): + # The GEMM base carries one cluster shape, and the fallback is the shape it + # must be sized against: every launch is guaranteed to be able to form the + # fallback cluster, so the cluster layout and multicast counts derived from + # it stay valid on both dispatch paths. + super().__init__( + acc_dtype=acc_dtype, + use_2cta_instrs=use_2cta_instrs, + mma_tiler_mn=mma_tiler_mn, + cluster_shape_mn=fallback_cluster_shape_mn, + use_tma_store=True, # Conv always uses im2col TMA store + swizzle_size=swizzle_size, + raster_along=raster_along, + ) + + # Both shapes are kept: the megakernel dispatches on which one the launch + # actually formed. + self.preferred_cluster_shape_mn = preferred_cluster_shape_mn + self.fallback_cluster_shape_mn = fallback_cluster_shape_mn + + def _setup_conv_input_attrs(self, a, b, d): + """Validate and set input-dependent attributes. + + Sets a_dtype, b_dtype, d_dtype, a_major_mode, b_major_mode, d_layout. + + :param a: Input tensor A - (N, D, H, W, C) layout + :param b: Filter tensor B - (K, T, R, S, C) layout + :param d: Output tensor D - (N, Z, P, Q, K) layout + """ + self.a_dtype: Type[cutlass.Numeric] = a.element_type + self.b_dtype: Type[cutlass.Numeric] = b.element_type + self.d_dtype: Type[cutlass.Numeric] = d.element_type + # Only C major accepted + if cutlass.const_expr(a.leading_dim != 4): + raise RuntimeError("The layout of a is not supported") + if cutlass.const_expr(b.leading_dim != 4): + raise RuntimeError("The layout of b is not supported") + if cutlass.const_expr(d.leading_dim != 4): + raise RuntimeError("The layout of d is not supported") + self.a_major_mode = cute.nvgpu.OperandMajorMode.K + self.b_major_mode = cute.nvgpu.OperandMajorMode.K + self.d_layout = ( + cutlass.tensor_utils.LayoutEnum.ROW_MAJOR + ) # K dimension contiguous + + # Check if input data types are compatible with MMA instruction + if cutlass.const_expr(self.a_dtype != self.b_dtype): + raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}") + + def _make_a_im2col_tma_atoms( + self, + a_tensor: cute.Tensor, + b_tensor: cute.Tensor, + tiled_mma: cute.TiledMma, + upper_pad_op: Tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32], + lower_pad_op: Tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32], + stride_op: Tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32], + dil_op: Tuple[cutlass.Int32, cutlass.Int32, cutlass.Int32], + ) -> Tuple[cute.CopyAtom, cute.Tensor, cute.CopyAtom, cute.Tensor]: + """Make A's preferred and fallback im2col TMA load atoms. + + A is the operand the im2col descriptor walks: it is viewed as + ((W, H, D, N), C) so the descriptor's box slides over the spatial modes, + and its corners come from the runtime pad/stride/dilation operands and the + filter tensor's own extents, which is what lets one cubin serve any + geometry. The two atoms differ only in the cluster shape they multicast + over, and each drops the multicast op when its cluster does not multicast A. + + :param a_tensor: Input tensor A - (N, D, H, W, C) layout + :param b_tensor: Filter tensor B - (K, T, R, S, C) layout + :param tiled_mma: Tiled MMA configuration + :param upper_pad_op: Upper padding (D, H, W) as runtime cutlass.Int32 tuple + :param lower_pad_op: Lower padding (D, H, W) as runtime cutlass.Int32 tuple + :param stride_op: Convolution stride (D, H, W) as runtime cutlass.Int32 tuple + :param dil_op: Dilation (D, H, W) as runtime cutlass.Int32 tuple + + :return: (tma_atom_a_preferred, tma_tensor_a_preferred, + tma_atom_a_fallback, tma_tensor_a_fallback) + """ + # Create 2-mode hierarchical tensor layout: (N, D, H, W, C) -> ((W, H, D, N), C) + mA = cute.make_tensor( + a_tensor.iterator, cute.select(a_tensor.layout, mode=[3, 2, 1, 0, 4]) + ) + mA = cute.group_modes(mA, begin=0, end=4) + + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0)) + a_internal_type = ( + cutlass.TFloat32 if mA.element_type is cutlass.Float32 else None + ) + + # Filter T/R/S come from the filter tensor (K, T, R, S, C), whose dynamic + # layout makes the extents runtime Int32, so one cubin serves any T/R/S. + rt_filter_trs = (b_tensor.shape[1], b_tensor.shape[2], b_tensor.shape[3]) + + # --- A preferred --- + a_copy_atom_preferred = ( + cpasync.CopyBulkTensorIm2ColG2SMulticastOp(cta_group=self.cta_group) + if self.is_preferred_a_mcast + else cpasync.CopyBulkTensorIm2ColG2SOp(cta_group=self.cta_group) + ) + tma_atom_a_preferred, tma_tensor_a_preferred = ( + cute.nvgpu.make_im2col_tma_atom_A( + a_copy_atom_preferred, + mA, + a_smem_layout, + self.mma_tiler, + tiled_mma, + rt_filter_trs, + upper_pad_op, + lower_pad_op, + stride_op, + dil_op, + self.preferred_cluster_layout_vmnk.shape, + internal_type=a_internal_type, + ) + ) + + # --- A fallback --- + a_copy_atom_fallback = ( + cpasync.CopyBulkTensorIm2ColG2SMulticastOp(cta_group=self.cta_group) + if self.is_fallback_a_mcast + else cpasync.CopyBulkTensorIm2ColG2SOp(cta_group=self.cta_group) + ) + tma_atom_a_fallback, tma_tensor_a_fallback = cute.nvgpu.make_im2col_tma_atom_A( + a_copy_atom_fallback, + mA, + a_smem_layout, + self.mma_tiler, + tiled_mma, + rt_filter_trs, + upper_pad_op, + lower_pad_op, + stride_op, + dil_op, + self.fallback_cluster_layout_vmnk.shape, + internal_type=a_internal_type, + ) + return ( + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + ) + + def _setup_conv_tma( + self, a, b, d, tiled_mma, upper_pad_op, lower_pad_op, stride_op, dil_op + ): + """Set up TMA atoms and tensors for im2col convolution with dual cluster shapes. + + Creates preferred and fallback TMA atoms for A and B tensors, and a single + TMA atom for D (im2col store is cluster-independent). + + The pad/stride/dilation operands feed the im2col A descriptor corners. + Threading them as runtime cutlass.Int32 lets one compiled cubin serve + any pad/stride/dilation config. + + :param a: Input tensor A - (N, D, H, W, C) layout + :param b: Filter tensor B - (K, T, R, S, C) layout + :param d: Output tensor D - (N, Z, P, Q, K) layout + :param tiled_mma: Tiled MMA configuration + :param upper_pad_op: Upper padding (D, H, W) as runtime cutlass.Int32 tuple + :param lower_pad_op: Lower padding (D, H, W) as runtime cutlass.Int32 tuple + :param stride_op: Convolution stride (D, H, W) as runtime cutlass.Int32 tuple + :param dil_op: Dilation (D, H, W) as runtime cutlass.Int32 tuple + :returns: (tma_atom_a_preferred, tma_tensor_a_preferred, + tma_atom_a_fallback, tma_tensor_a_fallback, + tma_atom_b_preferred, tma_tensor_b_preferred, + tma_atom_b_fallback, tma_tensor_b_fallback, + tma_atom_d, tma_tensor_d) + """ + atom_thr_size = cute.size(tiled_mma.thr_id.shape) + + ( + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + ) = self._make_a_im2col_tma_atoms( + a, b, tiled_mma, upper_pad_op, lower_pad_op, stride_op, dil_op + ) + + # --- B: tiled TMA load --- + b_smem_layout = cute.slice_(self.b_smem_layout_staged, (None, None, None, 0)) + # Change view of filter tensor from (K, T, R, S, C) to (K, (C, S, R, T)) + mB = cute.make_tensor(b.iterator, cute.select(b.layout, mode=[0, 4, 3, 2, 1])) + mB = cute.group_modes(mB, begin=1, end=5) + b_internal_type = ( + cutlass.TFloat32 if mB.element_type is cutlass.Float32 else None + ) + + # --- B preferred --- + b_op_preferred = sm100_utils.cluster_shape_to_tma_atom_B( + self.preferred_cluster_shape_mn, tiled_mma.thr_id + ) + tma_atom_b_preferred, tma_tensor_b_preferred = cute.nvgpu.make_tiled_tma_atom_B( + b_op_preferred, + mB, + b_smem_layout, + self.mma_tiler, + tiled_mma, + self.preferred_cluster_layout_vmnk.shape, + internal_type=b_internal_type, + ) + + # --- B fallback --- + b_op_fallback = sm100_utils.cluster_shape_to_tma_atom_B( + self.fallback_cluster_shape_mn, tiled_mma.thr_id + ) + tma_atom_b_fallback, tma_tensor_b_fallback = cute.nvgpu.make_tiled_tma_atom_B( + b_op_fallback, + mB, + b_smem_layout, + self.mma_tiler, + tiled_mma, + self.fallback_cluster_layout_vmnk.shape, + internal_type=b_internal_type, + ) + + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, None, 0)) + a_copy_size = cute.size_in_bytes(self.a_dtype, a_smem_layout) + b_copy_size = cute.size_in_bytes(self.b_dtype, b_smem_layout) + self.num_tma_load_bytes = (a_copy_size + b_copy_size) * atom_thr_size + # Response size is 4B * 4 elements + self.num_clc_response_bytes = 16 + + # --- D: im2col TMA store (cluster-independent) --- + # Change view of output tensor from (N, Z, P, Q, K) to ((Q, P, Z, N), K) + mD = cute.make_tensor(d.iterator, cute.select(d.layout, mode=[3, 2, 1, 0, 4])) + mD = cute.group_modes(mD, begin=0, end=4) + + epi_smem_layout = cute.slice_( + self.d_smem_layout_staged, (None, None, (None, 0)) + ) + + tma_atom_d, tma_tensor_d = cpasync.make_im2col_tma_atom( + cpasync.CopyBulkTensorIm2ColS2GOp(), + mD, + epi_smem_layout, + self.epi_tile, + ) + tma_tensor_d = cute.coalesce(tma_tensor_d, target_profile=(1, 1)) + + # Add dummy batch dimension to all tensors (GEMM expects batch dimension) + def add_dummy_batch_dimension(tensor): + new_layout = cute.append(tensor.layout, cute.make_layout(1)) + tensor = cute.make_tensor(tensor.iterator, new_layout) + return tensor + + tma_tensor_a_preferred = add_dummy_batch_dimension(tma_tensor_a_preferred) + tma_tensor_a_fallback = add_dummy_batch_dimension(tma_tensor_a_fallback) + tma_tensor_b_preferred = add_dummy_batch_dimension(tma_tensor_b_preferred) + tma_tensor_b_fallback = add_dummy_batch_dimension(tma_tensor_b_fallback) + tma_tensor_d = add_dummy_batch_dimension(tma_tensor_d) + + return ( + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + tma_atom_b_preferred, + tma_tensor_b_preferred, + tma_atom_b_fallback, + tma_tensor_b_fallback, + tma_atom_d, + tma_tensor_d, + ) + + def _setup_attributes(self): + """Set up configurations that are dependent on convolution inputs.""" + # Configure tiled mma + tiled_mma = self._create_tiled_mma() + + # Compute mma/cluster/tile shapes + mma_inst_shape_k = cute.size(tiled_mma.shape_mnk, mode=[2]) + mma_inst_tile_k = 4 + self.mma_tiler = ( + self.mma_tiler[0], + self.mma_tiler[1], + (mma_inst_shape_k * mma_inst_tile_k,), + ) + self.cta_tile_shape_mnk = ( + self.mma_tiler[0] // cute.size(tiled_mma.thr_id.shape), + self.mma_tiler[1], + self.mma_tiler[2], + ) + + # Compute cluster layout + self.cluster_layout_vmnk = cute.tiled_divide( + cute.make_layout((*self.cluster_shape_mn, 1)), + (tiled_mma.thr_id.shape,), + ) + + # Compute number of multicast CTAs for A/B + self.num_mcast_ctas_a = cute.size(self.cluster_layout_vmnk.shape[2]) + self.num_mcast_ctas_b = cute.size(self.cluster_layout_vmnk.shape[1]) + self.is_a_mcast = self.num_mcast_ctas_a > 1 + self.is_b_mcast = self.num_mcast_ctas_b > 1 + + # Compute epilogue subtile + self.epi_tile = utils.sm100.compute_epilogue_tile_shape( + self.cta_tile_shape_mnk, + self.use_2cta_instrs, + self.d_layout, + self.d_dtype, + ) + + d_smem_layout = utils.sm100.make_smem_layout_epi( + self.d_dtype, self.d_layout, self.epi_tile, 1 + ) + + self.smem_capacity = cutlass.memory.get_smem_capacity_in_bytes() + + # Setup A/B/C stage count in shared memory and ACC stage count in tensor memory + self.num_acc_stage, self.num_ab_stage, self.num_d_stage = _compute_stages( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.b_dtype, + self.d_dtype, + self.smem_capacity, + self.occupancy, + self.use_tma_store, + d_smem_layout, + ) + + # Setup clc stage by default + self.num_clc_stage = 1 + assert self.num_clc_stage == 1, "Only single-stage CLC pipeline is supported" + + # Compute A/B/D shared memory layout + ( + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.d_smem_layout_staged, + ) = self._make_smem_layouts( + tiled_mma, + self.mma_tiler, + self.a_dtype, + self.b_dtype, + self.d_dtype, + self.d_layout, + self.epi_tile, + self.num_ab_stage, + self.num_d_stage, + ) + + # Compute the number of tensor memory allocation columns + self.num_tmem_alloc_cols = self._compute_num_tmem_alloc_cols( + tiled_mma, self.mma_tiler, self.num_acc_stage, self.arch + ) + + # Compute preferred cluster layout alongside the fallback one set above. + self.preferred_cluster_layout_vmnk = cute.tiled_divide( + cute.make_layout((*self.preferred_cluster_shape_mn, 1)), + (tiled_mma.thr_id.shape,), + ) + # Calculate multicast CTA counts for preferred cluster + self.num_preferred_mcast_ctas_a = cute.size( + self.preferred_cluster_layout_vmnk.shape[2] + ) + self.num_preferred_mcast_ctas_b = cute.size( + self.preferred_cluster_layout_vmnk.shape[1] + ) + self.is_preferred_a_mcast = self.num_preferred_mcast_ctas_a > 1 + self.is_preferred_b_mcast = self.num_preferred_mcast_ctas_b > 1 + + self.fallback_cluster_layout_vmnk = self.cluster_layout_vmnk + self.is_fallback_a_mcast = self.is_a_mcast + self.is_fallback_b_mcast = self.is_b_mcast + + @cute.jit + def __call__( + self, + a: cute.Tensor, + b: cute.Tensor, + d: cute.Tensor, + rt_upper_pad_d: cutlass.Int32, + rt_upper_pad_h: cutlass.Int32, + rt_upper_pad_w: cutlass.Int32, + rt_lower_pad_d: cutlass.Int32, + rt_lower_pad_h: cutlass.Int32, + rt_lower_pad_w: cutlass.Int32, + rt_stride_d: cutlass.Int32, + rt_stride_h: cutlass.Int32, + rt_stride_w: cutlass.Int32, + rt_dil_d: cutlass.Int32, + rt_dil_h: cutlass.Int32, + rt_dil_w: cutlass.Int32, + stream: cuda.CUstream, + bias_tensor: Optional[cute.Tensor] = None, + epilogue_op: cutlass.Constexpr = lambda x: x, + ): + """Execute the persistent convolution operation with dynamic preferred cluster scheduling. + + :param a: Input tensor A - (N, D, H, W, C) layout + :param b: Filter tensor B - (K, T, R, S, C) layout + :param d: Output tensor D - (N, Z, P, Q, K) layout + :param rt_upper_pad_d/h/w: Runtime upper padding (D, H, W) as Int32, so one + compiled cubin serves any padding config without recompilation + :param rt_lower_pad_d/h/w: Runtime lower padding (D, H, W) as Int32 + :param rt_stride_d/h/w: Runtime convolution stride (D, H, W) as Int32 + :param rt_dil_d/h/w: Runtime dilation (D, H, W) as Int32 + :param stream: CUDA stream for asynchronous execution + :param bias_tensor: Optional length-K device tensor holding a + per-output-channel bias added in FP32 as D = epilogue_op(acc + bias). + Pass None for a cubin without the bias path. + :param epilogue_op: Optional elementwise lambda function to apply to the output tensor + """ + self._setup_conv_input_attrs(a, b, d) + # The staged bias row is sized from the bias tensor's dtype while the + # epilogue reads it as the output type, so a mismatch would overrun the + # smem buffer. + self.has_bias: bool = bias_tensor is not None + if cutlass.const_expr( + self.has_bias and bias_tensor.element_type is not self.d_dtype + ): + raise TypeError( + f"bias dtype ({bias_tensor.element_type}) must match output " + f"dtype ({self.d_dtype})" + ) + self._setup_attributes() + + # Pack the runtime pad/stride/dilation scalars into (D, H, W) tuples that + # feed the im2col A descriptor corners. Keeping them as runtime Int32 lets + # a single compiled cubin run any pad/stride/dilation configuration. + upper_pad_op = (rt_upper_pad_d, rt_upper_pad_h, rt_upper_pad_w) + lower_pad_op = (rt_lower_pad_d, rt_lower_pad_h, rt_lower_pad_w) + stride_op = (rt_stride_d, rt_stride_h, rt_stride_w) + dil_op = (rt_dil_d, rt_dil_h, rt_dil_w) + + tiled_mma = self._create_tiled_mma() + + ( + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + tma_atom_b_preferred, + tma_tensor_b_preferred, + tma_atom_b_fallback, + tma_tensor_b_fallback, + tma_atom_d, + tma_tensor_d, + ) = self._setup_conv_tma( + a, b, d, tiled_mma, upper_pad_op, lower_pad_op, stride_op, dil_op + ) + + # Compute grid size and scheduler params for both cluster shapes + self.fallback_tile_sched_params, _ = self._compute_grid( + tma_tensor_d, + self.cta_tile_shape_mnk, + self.fallback_cluster_shape_mn, + self.swizzle_size, + self.raster_along, + ) + self.preferred_tile_sched_params, preferred_grid = self._compute_grid( + tma_tensor_d, + self.cta_tile_shape_mnk, + self.preferred_cluster_shape_mn, + self.swizzle_size, + self.raster_along, + ) + + # Build mBias: the per-output-channel bias broadcast to mD's + # ((Q,P,Z,N), K, 1) = (M, N, L) profile. It varies along the output channel + # K (GEMM-N, stride 1) and repeats across every spatial position (GEMM-M + # modes carry stride 0). The epilogue stages one CTA tile's channels into + # shared memory and reads them back from there, so the M extent only has + # to span a CTA tile of output rows; the real output-M size arrives at the + # store site through the tile coordinate. + if cutlass.const_expr(self.has_bias): + mBias_layout = cute.make_layout( + ( + (self.cta_tile_shape_mnk[0], 1, 1, 1), + cute.size(d, mode=[4]), + 1, + ), + stride=((0, 0, 0, 0), 1, 0), + ) + mBias_mnl = cute.make_tensor(bias_tensor.iterator, mBias_layout) + else: + mBias_mnl = None + + # Launch the megakernel synchronously + self.kernel( + tiled_mma, + tma_atom_a_preferred, + tma_tensor_a_preferred, + tma_atom_a_fallback, + tma_tensor_a_fallback, + tma_atom_b_preferred, + tma_tensor_b_preferred, + tma_atom_b_fallback, + tma_tensor_b_fallback, + tma_atom_d, + tma_tensor_d, + self.preferred_cluster_layout_vmnk, + self.fallback_cluster_layout_vmnk, + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.d_smem_layout_staged, + self.epi_tile, + self.preferred_tile_sched_params, + self.fallback_tile_sched_params, + mBias_mnl, + epilogue_op, + ).launch( + grid=preferred_grid, + block=[self.threads_per_cta, 1, 1], + cluster=(*self.preferred_cluster_shape_mn, 1), + fallback_cluster=(*self.fallback_cluster_shape_mn, 1), + stream=stream, + smem_merge_branch_allocs=True, + ) + return + + def check_mma_tiler_and_cluster_shape(self): + """Check if mma tiler, fallback and preferred cluster shapes are valid.""" + # Call parent validation + super().check_mma_tiler_and_cluster_shape() + + # Validate preferred cluster shape + if self.preferred_cluster_shape_mn[0] % (2 if self.use_2cta_instrs else 1) != 0: + raise testing.CantImplementError( + f"Invalid preferred cluster shape M: {self.preferred_cluster_shape_mn[0]}" + ) + + is_power_of_2 = lambda x: x > 0 and (x & (x - 1)) == 0 + if ( + self.preferred_cluster_shape_mn[0] * self.preferred_cluster_shape_mn[1] > 16 + or self.preferred_cluster_shape_mn[0] <= 0 + or self.preferred_cluster_shape_mn[1] <= 0 + or not is_power_of_2(self.preferred_cluster_shape_mn[0]) + or not is_power_of_2(self.preferred_cluster_shape_mn[1]) + ): + raise testing.CantImplementError( + f"Invalid preferred cluster shape: {self.preferred_cluster_shape_mn}" + ) + + # Check preferred is multiple of fallback + if ( + self.preferred_cluster_shape_mn[0] % self.fallback_cluster_shape_mn[0] != 0 + or self.preferred_cluster_shape_mn[1] % self.fallback_cluster_shape_mn[1] + != 0 + ): + raise testing.CantImplementError( + f"Preferred cluster shape {self.preferred_cluster_shape_mn} must be " + f"integer multiple of fallback cluster shape {self.fallback_cluster_shape_mn}" + ) + + def check_preferred_cluster_problem_size(self, m: int, n: int) -> None: + """Require enough CTA tiles along each dimension to form one preferred cluster. + + The persistent scheduler tiles the problem into CTA tiles and then groups + them into clusters. A preferred cluster of (pref_m, pref_n) CTAs needs at + least that many tiles along the matching dimension, so below that the + preferred cluster shape is unusable and the config is rejected early. The + check is per dimension: a total tile count large enough for the whole + cluster still fails if one dimension alone is short. + + A trailing partial tile still occupies a full CTA in the grid, so tiles are + counted with a round-up rather than by measuring the problem extent against + the cluster's whole tile footprint. + + :param m: Implicit-GEMM M dimension (N*Z*P*Q) + :param n: Implicit-GEMM N dimension (K) + + :raises testing.CantImplementError: If either dimension has too few tiles + """ + # This runs before any input-dependent setup, so the CTA tile is derived + # from the MMA tiler here rather than read off the instance. + cta_tile_m = self.mma_tiler_mn[0] // (2 if self.use_2cta_instrs else 1) + cta_tile_n = self.mma_tiler_mn[1] + m_tiles = -(-m // cta_tile_m) + n_tiles = -(-n // cta_tile_n) + pref_m, pref_n = self.preferred_cluster_shape_mn + if m_tiles < pref_m or n_tiles < pref_n: + raise testing.CantImplementError( + f"CTA tile count ({m_tiles}, {n_tiles}) from problem ({m}, {n}) with " + f"cta_tile=({cta_tile_m}, {cta_tile_n}) cannot form one preferred " + f"cluster preferred_cluster_shape_mn={self.preferred_cluster_shape_mn}" + ) + + def can_implement( + self, + ncdhw: Tuple[int, int, int, int, int], + k: int, + ab_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], + filter_trs: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + dilation_dhw: Tuple[int, int, int], + ) -> bool: + """Determine if the given tensor configuration can be implemented by this kernel. + + The convolution geometry is passed in rather than held on the instance: + the kernel takes it as runtime Int32 operands and reads the filter extents + off the filter tensor, so one cubin serves any geometry and there is no + static copy to read here. + """ + try: + sm100_utils.check_int8_mma_arch(ab_dtype) + self.check_supported_dtypes(ab_dtype, ab_dtype, d_dtype) + # Validate fallback (held by the base as cluster_shape_mn), preferred + # shape, and their multiple relation + self.check_mma_tiler_and_cluster_shape() + _check_tensor_alignment(ncdhw[1], k, ab_dtype, d_dtype) + # Compute implicit GEMM M + z, p, q = compute_zpq( + ncdhw[2:], + filter_trs, + stride_dhw, + upper_padding_dhw, + lower_padding_dhw, + dilation_dhw, + ) + self.check_preferred_cluster_problem_size(ncdhw[0] * z * p * q, k) + _check_im2col_descriptor_limits( + filter_trs, + upper_padding_dhw, + lower_padding_dhw, + stride_dhw, + dilation_dhw, + ) + _check_swizzle_size( + ncdhw[0] * z * p * q, + k, + self.mma_tiler_mn, + self.use_2cta_instrs, + self.preferred_cluster_shape_mn, + self.fallback_cluster_shape_mn, + self.swizzle_size, + self.raster_along, + ) + except testing.CantImplementError as e: + print(e) + return False + return True + + # GPU device kernel - megakernel dispatcher for preferred/fallback cluster + @cute.kernel + def kernel( + self, + tiled_mma: cute.TiledMma, + tma_atom_a_preferred: cute.CopyAtom, + mA_mkl_preferred: cute.Tensor, + tma_atom_a_fallback: cute.CopyAtom, + mA_mkl_fallback: cute.Tensor, + tma_atom_b_preferred: cute.CopyAtom, + mB_nkl_preferred: cute.Tensor, + tma_atom_b_fallback: cute.CopyAtom, + mB_nkl_fallback: cute.Tensor, + tma_atom_d: Optional[cute.CopyAtom], + mD_mnl: cute.Tensor, + preferred_cluster_layout_vmnk: cute.Layout, + fallback_cluster_layout_vmnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + d_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout, None], + epi_tile: cute.Tile, + preferred_tile_sched_params: ClcDynamicPersistentTileSchedulerParams, + fallback_tile_sched_params: ClcDynamicPersistentTileSchedulerParams, + mBias_mnl: Optional[cute.Tensor], + epilogue_op: cutlass.Constexpr, + ): + """ + GPU device kernel entry point for kernel with preferred cluster shape. + """ + # Get cluster coordinates to determine if this is a preferred cluster + cbdim_x, cbdim_y, cbdim_z = cute.arch.block_in_cluster_dim() + is_preferred_cluster = ( + cbdim_x == self.preferred_cluster_shape_mn[0] + and cbdim_y == self.preferred_cluster_shape_mn[1] + and cbdim_z == 1 + ) + + # Megakernel approach: two mutually exclusive code branches, only one path runs per launch. + # smem_merge_branch_allocs=True at launch enables shared memory reuse between two paths. + if is_preferred_cluster: + self.cluster_specific_kernel( + tiled_mma, + tma_atom_a_preferred, + mA_mkl_preferred, + tma_atom_b_preferred, + mB_nkl_preferred, + tma_atom_d, + mD_mnl, + preferred_cluster_layout_vmnk, + a_smem_layout_staged, + b_smem_layout_staged, + d_smem_layout_staged, + epi_tile, + preferred_tile_sched_params, + mBias_mnl, + epilogue_op, + self.is_preferred_a_mcast, + self.is_preferred_b_mcast, + self.preferred_cluster_shape_mn, + ) + else: + self.cluster_specific_kernel( + tiled_mma, + tma_atom_a_fallback, + mA_mkl_fallback, + tma_atom_b_fallback, + mB_nkl_fallback, + tma_atom_d, + mD_mnl, + fallback_cluster_layout_vmnk, + a_smem_layout_staged, + b_smem_layout_staged, + d_smem_layout_staged, + epi_tile, + fallback_tile_sched_params, + mBias_mnl, + epilogue_op, + self.is_fallback_a_mcast, + self.is_fallback_b_mcast, + self.fallback_cluster_shape_mn, + ) + + @staticmethod + def _make_smem_layouts( + tiled_mma: cute.TiledMma, + mma_tiler: Tuple[int, int, int], + a_dtype: Type[cutlass.Numeric], + b_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], + d_layout: cutlass.tensor_utils.LayoutEnum, + epi_tile: cute.Tile, + ab_stage: int, + d_stage: int, + ) -> Tuple[cute.ComposedLayout, cute.ComposedLayout, cute.ComposedLayout]: + """Build the staged shared memory layouts for A, B and the epilogue. + + :param tiled_mma: The tiled MMA the operands feed + :param mma_tiler: The MMA tiler shape (M, N, K) + :param a_dtype: Element type of A + :param b_dtype: Element type of B + :param d_dtype: Element type of D + :param d_layout: Layout enum of D + :param epi_tile: The epilogue tiler + :param ab_stage: Number of A/B pipeline stages + :param d_stage: Number of epilogue pipeline stages + + :return: The A, B and epilogue layouts, each staged + """ + a_smem_layout_staged = utils.sm100.make_smem_layout_a( + tiled_mma, mma_tiler, a_dtype, ab_stage + ) + b_smem_layout_staged = utils.sm100.make_smem_layout_b( + tiled_mma, mma_tiler, b_dtype, ab_stage + ) + d_smem_layout_staged = utils.sm100.make_smem_layout_epi( + d_dtype, d_layout, epi_tile, d_stage + ) + return a_smem_layout_staged, b_smem_layout_staged, d_smem_layout_staged + + def _epilogue_tmem_copy_and_partition( + self, + tidx: cutlass.Int32, + tAcc: cute.Tensor, + tCgD: cute.Tensor, + epi_tile: cute.Tile, + use_2cta_instrs: Union[cutlass.Boolean, bool], + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """Make tiledCopy for tensor memory load, then use it to partition tensor + memory (source) and register array (destination). + + :param tidx: The thread index in epilogue warp groups + :param tAcc: The accumulator tensor to be copied and partitioned + :param tCgD: The global tensor D to be copied and partitioned + :param epi_tile: The epilogue tiler + :param use_2cta_instrs: Whether use_2cta_instrs is enabled + + :return: A tuple containing (tiled_copy_t2r, tTR_tAcc, tTR_rAcc) where: + - tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + - tTR_tAcc: The partitioned accumulator tensor + - tTR_rAcc: The accumulated tensor in register used to hold t2r results + """ + copy_atom_t2r = sm100_utils.get_tmem_load_op( + self.cta_tile_shape_mnk, + self.d_layout, + self.d_dtype, + self.acc_dtype, + epi_tile, + use_2cta_instrs, + ) + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, STAGE) + tAcc_epi = cute.flat_divide(tAcc, epi_tile) + # (EPI_TILE_M, EPI_TILE_N) + tiled_copy_t2r = tcgen05.make_tmem_copy( + copy_atom_t2r, tAcc_epi[(None, None, 0, 0, 0)] + ) + + thr_copy_t2r = tiled_copy_t2r.get_slice(tidx) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_M, STAGE) + tTR_tAcc = thr_copy_t2r.partition_S(tAcc_epi) + + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL) + tCgD_epi = cute.flat_divide(tCgD, epi_tile) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N, RestM, RestN, RestL) + tTR_gD = thr_copy_t2r.partition_D(tCgD_epi) + # (T2R, T2R_M, T2R_N) + tTR_rAcc = cute.make_rmem_tensor( + tTR_gD[(None, None, None, 0, 0, 0, 0, 0)].shape, self.acc_dtype + ) + return tiled_copy_t2r, tTR_tAcc, tTR_rAcc + + def _epilogue_smem_copy_and_partition( + self, + tiled_copy_t2r: cute.TiledCopy, + tTR_rD: cute.Tensor, + tidx: cutlass.Int32, + sD: cute.Tensor, + ) -> Tuple[cute.TiledCopy, cute.Tensor, cute.Tensor]: + """Make tiledCopy for shared memory store, then use it to partition register + array (source) and shared memory (destination). + + :param tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + :param tTR_rD: The partitioned accumulator tensor + :param tidx: The thread index in epilogue warp groups + :param sD: The shared memory tensor to be copied and partitioned + + :return: A tuple containing (tiled_copy_r2s, tRS_rD, tRS_sD) where: + - tiled_copy_r2s: The tiled copy operation for register to smem copy(r2s) + - tRS_rD: The partitioned tensor D (register source) + - tRS_sD: The partitioned tensor D (smem destination) + """ + copy_atom_r2s = sm100_utils.get_smem_store_op( + self.d_layout, self.d_dtype, self.acc_dtype, tiled_copy_t2r + ) + tiled_copy_r2s = cute.make_tiled_copy_D(copy_atom_r2s, tiled_copy_t2r) + # (R2S, R2S_M, R2S_N, PIPE_D) + thr_copy_r2s = tiled_copy_r2s.get_slice(tidx) + tRS_sD = thr_copy_r2s.partition_D(sD) + # (R2S, R2S_M, R2S_N) + tRS_rD = tiled_copy_r2s.retile(tTR_rD) + return tiled_copy_r2s, tRS_rD, tRS_sD + + def _epilogue_bias_smem_partition( + self, + tiled_copy_t2r: cute.TiledCopy, + tidx: cutlass.Int32, + sBias: cute.Tensor, + epi_tile: cute.Tile, + ) -> cute.Tensor: + """Partition the CTA-staged bias row for the epilogue's shared memory reads. + + The row holds this CTA tile's output channels, so it is viewed as a CTA + tile whose N axis is the channel (stride 1) and whose M axis repeats it + (stride 0), then partitioned through the same tmem->rmem thread layout as + the accumulator. Rooting that partition at the row itself is what puts a + thread's own tile-local N offset on the shared memory pointer: whenever + the CTA tile has fewer M rows than the epilogue has threads, the thread + layout also spans N, and a partition rooted anywhere else hands every + such thread the columns of the thread at N offset zero. + + :param tiled_copy_t2r: The tiled copy operation for tmem to register copy(t2r) + :param tidx: The thread index in epilogue warp groups + :param sBias: The staged bias row, one value per output channel of this CTA tile + :param epi_tile: The epilogue tiler + + :return: The staged row partitioned as (T2R, T2R_M, T2R_N, SUBTILE), with + the subtiles grouped to index alongside the accumulator's + """ + # (cta_tile_m, cta_tile_n) - one channel per N column, repeated over M + sBias_mn = cute.make_tensor( + sBias.iterator, + cute.make_layout( + (self.cta_tile_shape_mnk[0], self.cta_tile_shape_mnk[1]), + stride=(0, 1), + ), + ) + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N) + sBias_epi = cute.flat_divide(sBias_mn, epi_tile) + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N) + tTR_sBias = tiled_copy_t2r.get_slice(tidx).partition_D(sBias_epi) + return cute.group_modes(tTR_sBias, 3, 5) + + @cute.jit() + def cluster_specific_kernel( + self, + tiled_mma: cute.TiledMma, + tma_atom_a: cute.CopyAtom, + mA_mkl: cute.Tensor, + tma_atom_b: cute.CopyAtom, + mB_nkl: cute.Tensor, + tma_atom_d: Optional[cute.CopyAtom], + mD_mnl: cute.Tensor, + cluster_layout_vmnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + d_smem_layout_staged: Union[cute.Layout, cute.ComposedLayout, None], + epi_tile: cute.Tile, + tile_sched_params: ClcDynamicPersistentTileSchedulerParams, + mBias_mnl: Optional[cute.Tensor], + epilogue_op: cutlass.Constexpr, + effective_is_a_mcast: bool, + effective_is_b_mcast: bool, + cluster_shape: Tuple[int, int], + ): + """ + GPU device kernel performing the CLC dynamic persistent convolution computation. + """ + warp_idx = cute.arch.warp_idx() + warp_idx = cute.arch.make_warp_uniform(warp_idx) + + # + # Prefetch tma desc + # + if warp_idx == self.tma_warp_id: + cpasync.prefetch_descriptor(tma_atom_a) + cpasync.prefetch_descriptor(tma_atom_b) + cpasync.prefetch_descriptor(tma_atom_d) + + use_2cta_instrs = cute.size(tiled_mma.thr_id.shape) == 2 + + # + # Setup cta/thread coordinates + # + # Coords inside cluster + bidx, bidy, bidz = cute.arch.block_idx() + mma_tile_coord_v = bidx % cute.size(tiled_mma.thr_id.shape) + is_leader_cta = mma_tile_coord_v == 0 + cta_rank_in_cluster = cute.arch.make_warp_uniform( + cute.arch.block_idx_in_cluster() + ) + is_first_cta_in_cluster = cta_rank_in_cluster == 0 + block_in_cluster_coord_vmnk = cluster_layout_vmnk.get_flat_coord( + cta_rank_in_cluster + ) + # Coord inside cta + tidx, _, _ = cute.arch.thread_idx() + + # + # Alloc and init: a+b full/empty, accumulator full/empty, CLC, tensor memory dealloc barrier + # + # Define shared storage for kernel + @cute.struct + class SharedStorage: + ab_full_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_ab_stage * 2] + acc_full_mbar_ptr: cute.struct.MemRange[ + cutlass.Int64, self.num_acc_stage * 2 + ] + tmem_dealloc_mbar: cutlass.Int64 + tmem_holding_buf: cutlass.Int32 + clc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, 2] + clc_response: cute.struct.Align[ + cute.struct.MemRange[cutlass.Int32, 4], + 16, # Align bytes + ] + + smem = cutlass.memory.SmemAllocator() + storage = smem.allocate(SharedStorage) + + # Initialize mainloop ab_pipeline (barrier) and states + ab_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + ab_pipeline_consumer_group = pipeline.CooperativeGroup(pipeline.Agent.Warp) + ab_producer, ab_consumer = pipeline.PipelineTmaUmma.create( + barrier_storage=storage.ab_full_mbar_ptr.data_ptr(), + num_stages=self.num_ab_stage, + producer_group=ab_pipeline_producer_group, + consumer_group=ab_pipeline_consumer_group, + tx_count=self.num_tma_load_bytes, + cta_layout_vmnk=cluster_layout_vmnk, + enable_multicast_signaling=True, + defer_sync=True, + ).make_participants() + + # Initialize acc_pipeline (barrier) and states + acc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + num_acc_consumer_threads = len(self.epilogue_warp_id) * ( + 2 if use_2cta_instrs else 1 + ) + acc_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_acc_consumer_threads + ) + acc_pipeline = pipeline.PipelineUmmaAsync.create( + barrier_storage=storage.acc_full_mbar_ptr.data_ptr(), + num_stages=self.num_acc_stage, + producer_group=acc_pipeline_producer_group, + consumer_group=acc_pipeline_consumer_group, + cta_layout_vmnk=cluster_layout_vmnk, + defer_sync=True, + ) + + # Initialize clc_pipeline (barrier) and states + clc_pipeline_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + cluster_size = cute.size(cluster_shape) + num_clc_consumer_threads = 32 * ( + 1 + cluster_size * (1 + len(self.epilogue_warp_id) + 1) + ) + clc_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, num_clc_consumer_threads + ) + clc_pipeline = pipeline.PipelineClcFetchAsync.create( + barrier_storage=storage.clc_mbar_ptr.data_ptr(), + num_stages=self.num_clc_stage, + producer_group=clc_pipeline_producer_group, + consumer_group=clc_pipeline_consumer_group, + tx_count=self.num_clc_response_bytes, + cta_layout_vmnk=cluster_layout_vmnk, + defer_sync=True, + ) + + tmem_alloc_barrier = pipeline.NamedBarrier( + barrier_id=self.tmem_alloc_sync_bar_id, + num_threads=32 * len((self.mma_warp_id, *self.epilogue_warp_id)), + ) + # Tensor memory dealloc barrier init + tmem = cutlass.memory.TmemAllocator( + storage.tmem_holding_buf.ptr, + barrier_for_retrieve=tmem_alloc_barrier, + allocator_warp_id=self.epilogue_warp_id[0], + is_two_cta=use_2cta_instrs, + two_cta_tmem_dealloc_mbar_ptr=storage.tmem_dealloc_mbar.ptr, + ) + + # Cluster arrive after barrier init + pipeline_init_arrive(cluster_shape_mn=cluster_shape, is_relaxed=True) + + # Initial clc response pointer + clc_response_ptr = storage.clc_response.data_ptr() + + clc_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_clc_stage + ) + + # + # Setup smem tensor A/B/C + # + # (MMA, MMA_M, MMA_K, STAGE) + sA = smem.allocate_tensor( + element_type=self.a_dtype, + layout=a_smem_layout_staged.outer, + byte_alignment=128, + swizzle=a_smem_layout_staged.inner, + ) + # (MMA, MMA_N, MMA_K, STAGE) + sB = smem.allocate_tensor( + element_type=self.b_dtype, + layout=b_smem_layout_staged.outer, + byte_alignment=128, + swizzle=b_smem_layout_staged.inner, + ) + + # + # Compute multicast mask for A/B buffer full + # + a_full_mcast_mask = None + b_full_mcast_mask = None + if cutlass.const_expr( + effective_is_a_mcast or effective_is_b_mcast or use_2cta_instrs + ): + a_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=2 + ) + b_full_mcast_mask = cpasync.create_tma_multicast_mask( + cluster_layout_vmnk, block_in_cluster_coord_vmnk, mcast_mode=1 + ) + + # + # Local_tile partition global tensors + # + # (bM, bK, RestM, RestK, RestL) + gA_mkl = cute.local_tile( + mA_mkl, cute.slice_(self.mma_tiler, (None, 0, None)), (None, None, None) + ) + # (bN, bK, RestN, RestK, RestL) + gB_nkl = cute.local_tile( + mB_nkl, cute.slice_(self.mma_tiler, (0, None, None)), (None, None, None) + ) + # (bM, bN, RestM, RestN, RestL) + gD_mnl = cute.local_tile( + mD_mnl, cute.slice_(self.mma_tiler, (None, None, 0)), (None, None, None) + ) + k_tile_cnt = cute.size(gA_mkl, mode=[3]) + + # + # Partition global tensor for TiledMMA_A/B/C + # + thr_mma = tiled_mma.get_slice(mma_tile_coord_v) + # (MMA, MMA_M, MMA_K, RestM, RestK, RestL) + tCgA = thr_mma.partition_A(gA_mkl) + # (MMA, MMA_N, MMA_K, RestN, RestK, RestL) + tCgB = thr_mma.partition_B(gB_nkl) + # (MMA, MMA_M, MMA_N, RestM, RestN, RestL) + tCgD = thr_mma.partition_C(gD_mnl) + + # + # Partition global/shared tensor for TMA load A/B + # + # TMA load A partition_S/D + a_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_vmnk, (0, 0, None, 0)).shape + ) + # ((atom_v, rest_v), STAGE) + # ((atom_v, rest_v), RestM, RestK, RestL) + tAsA, tAgA = cpasync.tma_partition( + tma_atom_a, + block_in_cluster_coord_vmnk[2], + a_cta_layout, + cute.group_modes(sA, 0, 3), + cute.group_modes(tCgA, 0, 3), + ) + # TMA load B partition_S/D + b_cta_layout = cute.make_layout( + cute.slice_(cluster_layout_vmnk, (0, None, 0, 0)).shape + ) + # ((atom_v, rest_v), STAGE) + # ((atom_v, rest_v), RestM, RestK, RestL) + tBsB, tBgB = cpasync.tma_partition( + tma_atom_b, + block_in_cluster_coord_vmnk[1], + b_cta_layout, + cute.group_modes(sB, 0, 3), + cute.group_modes(tCgB, 0, 3), + ) + + # + # Partition shared/tensor memory tensor for TiledMMA_A/B/C + # + # (MMA, MMA_M, MMA_K, STAGE) + tCrA = tiled_mma.make_fragment_A(sA) + # (MMA, MMA_N, MMA_K, STAGE) + tCrB = tiled_mma.make_fragment_B(sB) + # (MMA, MMA_M, MMA_N) + acc_shape = tiled_mma.partition_shape_C(self.mma_tiler[:2]) + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_fake = tiled_mma.make_fragment_C( + cute.append(acc_shape, self.num_acc_stage) + ) + + # + # Cluster wait before tensor memory alloc + # + pipeline_init_wait(cluster_shape_mn=cluster_shape) + + # + # Construct the scheduler + # + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + + # + # Specialized TMA load warp + # + + if warp_idx == self.tma_warp_id: + # + # Persistent tile scheduling loop + # + while work_tile.is_valid_tile: + # Get tile coord from tile scheduler + cur_tile_coord = work_tile.tile_idx + mma_tile_coord_mnl = ( + cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape), + cur_tile_coord[1], + cur_tile_coord[2], + ) + + # + # Slice to per mma tile index + # + # ((atom_v, rest_v), RestK) + tAgA_slice = tAgA[ + (None, mma_tile_coord_mnl[0], None, mma_tile_coord_mnl[2]) + ] + # ((atom_v, rest_v), RestK) + tBgB_slice = tBgB[ + (None, mma_tile_coord_mnl[1], None, mma_tile_coord_mnl[2]) + ] + + # Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + ab_producer.reset() + peek_ab_empty_status = ab_producer.try_acquire() + + # + # Tma load loop + # + + # Set up coord iterator to avoid incurring idx2crd at runtime + # Permuted iteration order: K-tuple layout is (C,S,R,T) but we + # iterate in (S,R,T,C) colex order so S advances innermost. + # Tensor/TMA layouts are unchanged; we un-permute the coord at + # slice time so A and B see the original (C,S,R,T) coord. + k_shape_orig = cute.shape(tAgA_slice, mode=1) + k_shape_perm = ( + k_shape_orig[1], + k_shape_orig[2], + k_shape_orig[3], + k_shape_orig[0], + ) + coord_perm = cute.repeat_like(0, k_shape_perm) + for k_tile in cutlass.range(0, k_tile_cnt, 1, unroll=1): + # Conditionally wait for AB buffer empty + handle = ab_producer.acquire_and_advance(peek_ab_empty_status) + + # Un-permute (s,r,t,c) -> (c,s,r,t) for slice + coord_iter = ( + coord_perm[3], + coord_perm[0], + coord_perm[1], + coord_perm[2], + ) + + # TMA load A/B + cute.copy( + tma_atom_a, + tAgA_slice[(None, coord_iter)], + tAsA[(None, handle.index)], + tma_bar_ptr=handle.barrier, + mcast_mask=a_full_mcast_mask, + ) + cute.copy( + tma_atom_b, + tBgB_slice[(None, coord_iter)], + tBsB[(None, handle.index)], + tma_bar_ptr=handle.barrier, + mcast_mask=b_full_mcast_mask, + ) + + # Peek (try_wait) AB buffer empty for k_tile = prefetch_k_tile_cnt + k_tile + 1 + peek_ab_empty_status = cutlass.Boolean(1) + if handle.count + 1 < k_tile_cnt: + peek_ab_empty_status = ab_producer.try_acquire() + + coord_perm = cute.increment_coord(coord_perm, k_shape_perm) + + # + # Advance to next tile + # + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + # + # Wait A/B buffer empty + # + ab_producer.tail() + + # + # Specialized scheduler warp + # + if warp_idx == self.sched_warp_id and is_first_cta_in_cluster: + # + # Persistent tile scheduling loop + # + clc_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.ProducerConsumer, self.num_clc_stage + ) + + while work_tile.is_valid_tile: + # + # Advance to next tile + # + clc_pipeline.producer_acquire(clc_producer_state) + mbarrier_addr = clc_pipeline.producer_get_barrier(clc_producer_state) + tile_sched.advance_to_next_work(mbarrier_addr) + clc_producer_state.advance() + + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + clc_pipeline.producer_tail(clc_producer_state) + + # + # Specialized MMA warp + # + if warp_idx == self.mma_warp_id: + # + # Retrieving tensor memory ptr and make accumulator tensor + # + tmem.wait_for_alloc() + tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout) + + # + # Persistent tile scheduling loop + # + acc_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_acc_stage + ) + + while work_tile.is_valid_tile: + # Get tile coord from tile scheduler + cur_tile_coord = work_tile.tile_idx + mma_tile_coord_mnl = ( + cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape), + cur_tile_coord[1], + cur_tile_coord[2], + ) + + # Set tensor memory buffer for current tile + # (MMA, MMA_M, MMA_N) + tCtAcc = tCtAcc_base[(None, None, None, acc_producer_state.index)] + + # Peek (try_wait) AB buffer full for k_tile = 0 + ab_consumer.reset() + peek_ab_full_status = cutlass.Boolean(1) + if is_leader_cta: + peek_ab_full_status = ab_consumer.try_wait() + + # + # Wait for accumulator buffer empty + # + if is_leader_cta: + acc_pipeline.producer_acquire(acc_producer_state) + + # + # Mma mainloop + # + for k_tile in range(k_tile_cnt): + if is_leader_cta: + # Conditionally wait for AB buffer full + handle = ab_consumer.wait_and_advance(peek_ab_full_status) + + # tCtAcc += tCrA * tCrB + tiled_mma.set(tcgen05.Field.ACCUMULATE, k_tile != 0) + tile_crd = (None, None, None, handle.index) + cute.gemm( + tiled_mma, tCtAcc, tCrA[tile_crd], tCrB[tile_crd], tCtAcc + ) + + # Async arrive AB buffer empty + handle.release() + + # Peek (try_wait) AB buffer full for k_tile = k_tile + 1 + peek_ab_full_status = cutlass.Boolean(1) + if handle.count + 1 < k_tile_cnt: + peek_ab_full_status = ab_consumer.try_wait() + + # + # Async arrive accumulator buffer full + # + if is_leader_cta: + acc_pipeline.producer_commit(acc_producer_state) + acc_producer_state.advance() + + # + # Advance to next tile + # + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + # + # Wait for accumulator buffer empty + # + acc_pipeline.producer_tail(acc_producer_state) + + # (EPI_TILE_M, EPI_TILE_N, STAGE) + sD = smem.allocate_tensor( + element_type=self.d_dtype, + layout=d_smem_layout_staged.outer, + byte_alignment=128, + swizzle=d_smem_layout_staged.inner, + ) + + # (cta_tile_n,) staged bias row: one value per output channel in this + # CTA's N-tile, fetched once and read back by every epilogue thread. Only + # allocated on the bias path, so a no-bias cubin spends no smem on it. + if cutlass.const_expr(mBias_mnl is not None): + sBias = smem.allocate_tensor( + element_type=mBias_mnl.element_type, + layout=cute.make_layout(self.cta_tile_shape_mnk[1]), + byte_alignment=128, + ) + else: + sBias = None + + # Bias staging is a CTA-wide handoff: one warp's cp.async fills sBias and + # every epilogue thread reads its own columns back, so they meet here. + bias_sync_barrier = pipeline.NamedBarrier( + barrier_id=self.epilog_sync_bar_id, + num_threads=32 * len(self.epilogue_warp_id), + ) + + # + # Specialized epilogue warps + # + if warp_idx < self.mma_warp_id: + # + # Alloc tensor memory buffer + # + tmem.allocate(self.num_tmem_alloc_cols) + + # + # Retrieving tensor memory ptr and make accumulator tensor + # + tmem.wait_for_alloc() + tmem_ptr = tmem.retrieve_ptr(self.acc_dtype) + # (MMA, MMA_M, MMA_N, STAGE) + tCtAcc_base = cute.make_tensor(tmem_ptr, tCtAcc_fake.layout) + + # + # Persistent tile scheduling loop for epilogue + # + acc_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_acc_stage + ) + + assert tma_atom_d is not None and sD is not None + d_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + 32 * len(self.epilogue_warp_id), + ) + d_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.num_d_stage, producer_group=d_producer_group + ) + # Epilogue partitions, hoisted out of the tile loop: the tmem->rmem + # and rmem->smem tilings and the TMA store partition depend only on + # compile-time shapes. + tCgD_t = utils.gemm.sm100.transform_partitioned_tensor_layout(tCgD) + tCtAcc_t = utils.gemm.sm100.transform_partitioned_tensor_layout(tCtAcc_base) + ( + tiled_copy_t2r, + tTR_tAcc_base, + tTR_rAcc, + ) = self._epilogue_tmem_copy_and_partition( + tidx, tCtAcc_t, tCgD_t, epi_tile, self.use_2cta_instrs + ) + tTR_rD = cute.make_rmem_tensor(tTR_rAcc.shape, self.d_dtype) + ( + tiled_copy_r2s, + tRS_rD, + tRS_sD, + ) = self._epilogue_smem_copy_and_partition(tiled_copy_t2r, tTR_rD, tidx, sD) + # (EPI_TILE_M, EPI_TILE_N, EPI_M, EPI_N, RestM, RestN, RestL) + tCgD_epi = cute.flat_divide(tCgD_t, epi_tile) + # ((ATOM_V, REST_V), EPI_M, EPI_N[, RestM, RestN, RestL]) + bSG_sD, bSG_gD_partitioned = cpasync.tma_partition( + tma_atom_d, + 0, + cute.make_layout(1), + cute.group_modes(sD, 0, 2), + cute.group_modes(tCgD_epi, 0, 2), + ) + + # Bias read view: the staged row carries the same epi-tile and T2R + # structure as the accumulator, so each thread's bias fragment lines + # up with its acc fragment. It is tile invariant, since every tile + # restages the row in place. + if cutlass.const_expr(mBias_mnl is not None): + # (T2R, T2R_M, T2R_N, SUBTILE) + tTR_sBias = self._epilogue_bias_smem_partition( + tiled_copy_t2r, tidx, sBias, epi_tile + ) + # cp.async moves at least 32 bits, so each lane carries a + # 32-bit-wide vector of bias elements (2 bf16/fp16, 4 fp8). + bias_elems_per_copy = 32 // mBias_mnl.element_type.width + bias_g2s_atom = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(), + mBias_mnl.element_type, + num_bits_per_copy=32, + ) + else: + tTR_sBias = None + bias_elems_per_copy = None + bias_g2s_atom = None + + while work_tile.is_valid_tile: + # Get tile coord from tile scheduler + cur_tile_coord = work_tile.tile_idx + mma_tile_coord_mnl = ( + cur_tile_coord[0] // cute.size(tiled_mma.thr_id.shape), + cur_tile_coord[1], + cur_tile_coord[2], + ) + + num_tiles_executed = tile_sched.num_tiles_executed + + # ((ATOM_V, REST_V), EPI_M, EPI_N) + bSG_gD = bSG_gD_partitioned[(None, None, None, *mma_tile_coord_mnl)] + + if cutlass.const_expr(mBias_mnl is not None): + # One cp.async per lane brings this tile's contiguous + # cta_tile_n bias values (channel stride 1) into sBias; the + # barrier below then lets every epilogue thread read its own + # columns back out of smem. A channel's value is fetched once + # and reused by every output row of the tile. + cta_n = self.cta_tile_shape_mnk[1] + n_base = mma_tile_coord_mnl[1] * cta_n + n_active = cta_n // bias_elems_per_copy + # (elems, lanes), column-major so lane t owns the contiguous + # block [t*elems, (t+1)*elems). + row_layout = cute.make_layout( + (bias_elems_per_copy, n_active), + stride=(1, bias_elems_per_copy), + ) + # cp.async needs 32-bit alignment on both ends; n_base is a + # multiple of cta_tile_n, so re-annotate to satisfy the atom. + gBias_row = cute.make_tensor( + (mBias_mnl.iterator + n_base).align(min_align=4), + row_layout, + ) + sBias_tiled = cute.make_tensor( + sBias.iterator.align(min_align=4), row_layout + ) + # A CTA N-tile rounds up to cta_tile_n, but the channel count + # need not divide it, so tail lanes address columns with no + # backing storage. Predicate per lane: cp.async writes zero + # when the predicate is false, and 16-byte channel alignment + # keeps every 32-bit vector wholly in or out of bounds. The + # zero tail is only read back for overhang the TMA store + # clamps away. + # One pass covers as many columns as there are epilogue + # threads. A 32-bit vector holds a single FP32 element, so a + # cta_tile_n past that thread count needs more than one pass + # -- without them the columns past the first pass would be + # read back out of shared memory that nothing wrote. + num_bias_threads = 32 * len(self.epilogue_warp_id) + for lane_base in cutlass.range_constexpr( + 0, n_active, num_bias_threads + ): + lane = lane_base + tidx + if lane < n_active: + bias_pred = cute.make_rmem_tensor( + cute.make_layout((1,)), cutlass.Boolean + ) + bias_pred[0] = cutlass.Boolean( + n_base + lane * bias_elems_per_copy < mD_mnl.shape[1] + ) + cute.copy_atom_call( + bias_g2s_atom, + gBias_row[(None, lane)], + sBias_tiled[(None, lane)], + pred=bias_pred, + ) + cute.arch.cp_async_commit_group() + cute.arch.cp_async_wait_group(0) + bias_sync_barrier.arrive_and_wait() + + # (T2R, T2R_M, T2R_N, EPI_M, EPI_N) + tTR_tAcc = tTR_tAcc_base[ + (None, None, None, None, None, acc_consumer_state.index) + ] + + acc_pipeline.consumer_wait(acc_consumer_state) + + tTR_tAcc = cute.group_modes(tTR_tAcc, 3, cute.rank(tTR_tAcc)) + bSG_gD = cute.group_modes(bSG_gD, 1, cute.rank(bSG_gD)) + + subtile_cnt = cute.size(tTR_tAcc.shape, mode=[3]) + num_prev_subtiles = num_tiles_executed * subtile_cnt + for subtile_idx in cutlass.range_constexpr(subtile_cnt): + # Load the accumulator subtile from tensor memory + tTR_tAcc_mn = tTR_tAcc[(None, None, None, subtile_idx)] + cute.copy(tiled_copy_t2r, tTR_tAcc_mn, tTR_rAcc) + + # Add the bias on the FP32 accumulator, then narrow and + # activate: D = epilogue_op(acc + bias). + if cutlass.const_expr(mBias_mnl is not None): + tTR_rBias = cute.make_rmem_tensor( + tTR_rAcc.shape, mBias_mnl.element_type + ) + cute.autovec_copy( + tTR_sBias[(None, None, None, subtile_idx)], tTR_rBias + ) + tTR_rAcc.store( + tTR_rAcc.load() + tTR_rBias.load().to(self.acc_dtype) + ) + acc_vec = tiled_copy_r2s.retile(tTR_rAcc).load() + acc_vec = acc_vec.to(self.d_dtype) + acc_vec = epilogue_op(acc_vec) + tRS_rD.store(acc_vec) + + # Store the subtile to shared memory + d_buffer = (num_prev_subtiles + subtile_idx) % self.num_d_stage + cute.copy( + tiled_copy_r2s, tRS_rD, tRS_sD[(None, None, None, d_buffer)] + ) + cute.arch.fence_proxy("async.shared", space="cta") + bias_sync_barrier.arrive_and_wait() + + # TMA store the subtile to global memory + if warp_idx == self.epilogue_warp_id[0]: + cute.copy( + tma_atom_d, + bSG_sD[(None, d_buffer)], + bSG_gD[(None, subtile_idx)], + ) + d_pipeline.producer_commit() + d_pipeline.producer_acquire() + bias_sync_barrier.arrive_and_wait() + + bias_sync_barrier.arrive_and_wait() + + with cute.arch.elect_one(): + acc_pipeline.consumer_release(acc_consumer_state) + acc_consumer_state.advance() + # + # Advance to next tile + # + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + + # Wait for C store complete + d_pipeline.producer_tail() + + # + # Dealloc the tensor memory buffer + # + tmem.relinquish_alloc_permit() + tmem.free(tmem_ptr) + + +def _rt_conv_scalars( + *, + upper_pad: Tuple[int, int, int], + lower_pad: Tuple[int, int, int], + stride: Tuple[int, int, int], + dil: Tuple[int, int, int], +) -> list: + """Box the runtime conv geometry as cutlass.Int32 in the kernel's argument order. + + The order (upper pad, lower pad, stride, dilation, each D/H/W) must match the + kernel __call__ signature; keyword-only parameters keep the four groups from + being swapped at a call site. + """ + return [ + cutlass.Int32(v) for group in (upper_pad, lower_pad, stride, dil) for v in group + ] + + +# Compile-time epilogue activations, folded into the cubin as a Constexpr op +# (one activation per cubin). Each entry pairs the device-side op applied to the +# output fragment with the torch op used to build the reference. +EPILOGUE_ACTIVATIONS = { + "identity": { + "device": lambda x: x, + "ref": lambda x: x, + }, + "relu": { + "device": lambda x: cute.where(x > 0, x, cute.full_like(x, 0)), + "ref": torch.nn.functional.relu, + }, +} + + +@lru_cache(maxsize=1) +def compile_conv( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + input: cute.Tensor, + filter: cute.Tensor, + output: cute.Tensor, + acc_dtype: Type[cutlass.Numeric], + mma_tiler: Tuple[int, int] = (256, 256), + preferred_cluster_shape_mn: Tuple[int, int] = (2, 1), + fallback_cluster_shape_mn: Tuple[int, int] = (2, 1), + swizzle_size: int = 1, + raster_along: Literal["m", "n"] = "m", + use_2cta_instrs: bool = True, + upper_padding_dhw: Tuple[int, int, int] = (0, 0, 0), + lower_padding_dhw: Tuple[int, int, int] = (0, 0, 0), + stride_dhw: Tuple[int, int, int] = (1, 1, 1), + dilation_dhw: Tuple[int, int, int] = (1, 1, 1), + bias: Optional[cute.Tensor] = None, + epilogue_op: cutlass.Constexpr = lambda x: x, +): + """ + Compile a 3D convolution kernel with CLC dynamic preferred cluster scheduling. + + :param ncdhw: Problem shape (N, C, D, H, W) + :param ktrs: Problem shape (K, T, R, S) + :param input: Input tensor (N, D, H, W, C) with C contiguous + :param filter: Filter tensor in KTRSC format (K, T, R, S, C) with C contiguous + :param output: Output tensor (N, Z, P, Q, K) with K contiguous + :param acc_dtype: Accumulator data type + :param mma_tiler: MMA tile shape (M, N) + :param preferred_cluster_shape_mn: Preferred cluster shape (M, N) + :param fallback_cluster_shape_mn: Fallback cluster shape (M, N) + :param swizzle_size: Swizzling size in the unit of cluster. 1 means no swizzle + :param raster_along: Rasterization order of clusters. Only used when swizzle_size > 1 + :param use_2cta_instrs: Whether to use 2CTA instructions + :param upper_padding_dhw: Upper padding (PadD, PadH, PadW) + :param lower_padding_dhw: Lower padding (PadD, PadH, PadW) + :param stride_dhw: Stride (Sd, Sh, Sw) + :param dilation_dhw: Dilation (DilD, DilH, DilW) + :param bias: Optional length-K per-output-channel bias tensor, or None for a + cubin without the bias path + :param epilogue_op: Epilogue operation + :return: Compiled kernel function + """ + from cutlass.cute.runtime import make_fake_stream + + # Create convolution kernel object + conv_op = Sm100PersistentDenseImplicitGemmFpropKernel( + acc_dtype, + use_2cta_instrs, + mma_tiler, + preferred_cluster_shape_mn, + fallback_cluster_shape_mn, + swizzle_size, + raster_along, + ) + + # Check if configuration can be implemented + can_implement = conv_op.can_implement( + ncdhw, + ktrs[0], + input.element_type, + output.element_type, + ktrs[1:], + upper_padding_dhw, + lower_padding_dhw, + stride_dhw, + dilation_dhw, + ) + if not can_implement: + raise testing.CantImplementError("The current config is invalid/unsupported.") + + stream = make_fake_stream() + # Box pad/stride/dilation as runtime Int32 so cute.compile lowers them to + # SSA values and keeps them out of the mangled kernel name; one cubin then + # serves any pad/stride/dilation without recompilation. + return cute.compile( + conv_op, + input, + filter, + output, + *_rt_conv_scalars( + upper_pad=upper_padding_dhw, + lower_pad=lower_padding_dhw, + stride=stride_dhw, + dil=dilation_dhw, + ), + stream, + bias, + epilogue_op, + ) + + +def compute_zpq( + dhw: Tuple[int, int, int], + trs: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], + dilation_dhw: Tuple[int, int, int], +) -> Tuple[int, int, int]: + """Compute output spatial dimensions Z, P, and Q with asymmetric padding.""" + D, H, W = dhw + T, R, S = trs + Sd, Sh, Sw = stride_dhw + UpperPadD, UpperPadH, UpperPadW = upper_padding_dhw + LowerPadD, LowerPadH, LowerPadW = lower_padding_dhw + DilD, DilH, DilW = dilation_dhw + Z = ((D + UpperPadD + LowerPadD - DilD * (T - 1) - 1) // Sd) + 1 + P = ((H + UpperPadH + LowerPadH - DilH * (R - 1) - 1) // Sh) + 1 + Q = ((W + UpperPadW + LowerPadW - DilW * (S - 1) - 1) // Sw) + 1 + return Z, P, Q + + +def create_cute_tensor( + source_f32_tensor: torch.Tensor, + dtype: Type[cutlass.Numeric], + leading_dim: int = None, +) -> Tuple[cute.Tensor, torch.Tensor]: + """Create a cute tensor with dynamic layout from a source f32 tensor. + + :param source_f32_tensor: Source f32 tensor + :type source_f32_tensor: torch.Tensor + :param dtype: Data type + :type dtype: Type[cutlass.Numeric] + :param leading_dim: Leading dimension for dynamic layout + :type leading_dim: int + :return: Tuple of cute tensor and storage tensor + :rtype: Tuple[cute.Tensor, torch.Tensor] + """ + import cutlass.torch as cutlass_torch + + storage_type = torch.int8 if is_fp8_dtype(dtype) else torch_dtype(dtype) + storage_tensor = source_f32_tensor.to(dtype=storage_type) + + cute_tensor = from_dlpack( + storage_tensor, assumed_align=16, force_tf32=dtype == cutlass.TFloat32 + ) + cute_tensor = cute_tensor.mark_layout_dynamic(leading_dim=leading_dim) + if is_fp8_dtype(dtype): + cute_tensor.element_type = dtype + cute_tensor = cutlass_torch.convert_cute_tensor( + source_f32_tensor, cute_tensor, dtype, is_dynamic_layout=True + ) + # Cast the underlying storage tensor to the correct dtype + storage_tensor = storage_tensor.view(dtype=torch_dtype(dtype)) + return cute_tensor, storage_tensor + + +def prepare_tensors( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + zpq: Tuple[int, int, int], + ab_dtype: Type[cutlass.Numeric], +): + """Prepare f32 tensors for 3D convolution. + + :param ncdhw: Input tensor shape (N, C, D, H, W) + :type ncdhw: Tuple[int, int, int, int, int] + :param ktrs: Filter tensor shape components (K, T, R, S) + :type ktrs: Tuple[int, int, int, int] + :param zpq: Output spatial dimensions (Z, P, Q) + :type zpq: Tuple[int, int, int] + :param ab_dtype: Data type for A/B input tensors + :type ab_dtype: Type[cutlass.Numeric] + :return: Tuple of input, filter, and output tensors + :rtype: Tuple[torch.Tensor, torch.Tensor, torch.Tensor] + """ + N, C, D, H, W = ncdhw + K, T, R, S = ktrs + Z, P, Q = zpq + + # Initialize with small random values for numerical stability + if ab_dtype == cutlass.Uint8: + input_range = (0, 2) + else: + input_range = (-1, 2) + input_tensor = torch.randint( + input_range[0], + input_range[1], + (N, D, H, W, C), + dtype=torch.float32, + device="cuda", + ) + filter_tensor = torch.randint( + input_range[0], + input_range[1], + (K, T, R, S, C), + dtype=torch.float32, + device="cuda", + ) + output_tensor = torch.empty((N, Z, P, Q, K), dtype=torch.float32, device="cuda") + + return input_tensor, filter_tensor, output_tensor + + +def run( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + stride_dhw: Tuple[int, int, int] = (1, 1, 1), + upper_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + lower_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + dil_dhw: Tuple[int, int, int] = (1, 1, 1), + ab_dtype: Type[cutlass.Numeric] = cutlass.Float16, + d_dtype: Type[cutlass.Numeric] = cutlass.Float16, + acc_dtype: Type[cutlass.Numeric] = cutlass.Float32, + mma_tiler_mn: Tuple[int, int] = (128, 128), + preferred_cluster_shape_mn: Tuple[int, int] = (1, 1), + fallback_cluster_shape_mn: Tuple[int, int] = (1, 1), + swizzle_size: int = 1, + raster_along: Literal["m", "n"] = "m", + use_2cta_instrs: bool = False, + tolerance: float = 1e-02, + warmup_iterations: int = 0, + iterations: int = 1, + use_cold_l2: bool = False, + skip_ref_check: bool = False, + use_bias: bool = False, + activation: str = "identity", + **kwargs, +): + """Run 3D convolution and compare against PyTorch reference. + + The filter tensor uses native KTRSC layout (K, T, R, S, C) with C contiguous. + + :param ncdhw: Input tensor shape (N, C, D, H, W) + :param ktrs: Filter tensor shape components (K, T, R, S) + :param stride_dhw: Stride (Sd, Sh, Sw) + :param upper_pad_dhw: Upper padding (PadD, PadH, PadW) + :param lower_pad_dhw: Lower padding (PadD, PadH, PadW) + :param dil_dhw: Dilation (DilD, DilH, DilW) + :param ab_dtype: Data type for A/B input tensors + :param d_dtype: Data type for output tensor D + :param acc_dtype: Accumulator data type + :param mma_tiler_mn: MMA tiler shape + :param preferred_cluster_shape_mn: Preferred cluster shape (M, N) + :param fallback_cluster_shape_mn: Fallback cluster shape (M, N) + :param swizzle_size: Swizzling size in the unit of cluster. 1 means no swizzle + :param raster_along: Rasterization order of clusters. Only used when swizzle_size > 1 + :param use_2cta_instrs: Whether to use 2CTA instructions + :param tolerance: Tolerance for result comparison + :param warmup_iterations: Number of warmup iterations + :param iterations: Number of benchmark iterations + :param use_cold_l2: Whether to flush L2 cache between iterations + :param skip_ref_check: Whether to skip reference checking + :param use_bias: Add a per-output-channel bias in the output dtype as + D = activation(acc + bias); bias and no-bias are separate cubins + :param activation: Compile-time epilogue activation, one per cubin + """ + if not torch.cuda.is_available(): + raise RuntimeError("GPU is required to run this example!") + + N, C, D, H, W = ncdhw + K, T, R, S = ktrs + + Z, P, Q = compute_zpq( + (D, H, W), + (T, R, S), + stride_dhw, + upper_pad_dhw, + lower_pad_dhw, + dil_dhw, + ) + + print("Running Blackwell 3D Convolution test with:") + print(f" Input shape (N, C, D, H, W): {ncdhw}") + print(f" Filter shape (K, C, T, R, S): ({K}, {C}, {T}, {R}, {S})") + print(f" Output shape (N, K, Z, P, Q): ({N}, {K}, {Z}, {P}, {Q})") + print(f" Stride (Sd, Sh, Sw): {stride_dhw}") + print(f" Upper padding (PadD, PadH, PadW): {upper_pad_dhw}") + print(f" Lower padding (PadD, PadH, PadW): {lower_pad_dhw}") + print(f" Dilation (DilD, DilH, DilW): {dil_dhw}") + print(f" A/B data type: {ab_dtype}") + print(f" D data type: {d_dtype}") + print(f" Accumulator type: {acc_dtype}") + print(f" MMA tiler (M, N): {mma_tiler_mn}") + print(f" Preferred cluster shape (M, N): {preferred_cluster_shape_mn}") + print(f" Fallback cluster shape (M, N): {fallback_cluster_shape_mn}") + print(f" Swizzle size: {swizzle_size}") + print(f" Raster along: {raster_along}") + print(f" Use 2CTA instructions: {use_2cta_instrs}\n") + + # Create input and filter tensors + input_tensor, filter_tensor, output_tensor = prepare_tensors( + ncdhw, ktrs, (Z, P, Q), ab_dtype + ) + + # Prepare cute tensors + input_, input_storage = create_cute_tensor(input_tensor, ab_dtype, leading_dim=4) + filter_, filter_storage = create_cute_tensor(filter_tensor, ab_dtype, leading_dim=4) + output_, output_storage = create_cute_tensor(output_tensor, d_dtype, leading_dim=4) + + # One bias value per output channel, in the output dtype the epilogue reads + # it as. Small integers keep the fp8 case exact so the reference matches. + if use_bias: + bias_tensor = torch.randint( + -2, 3, (ktrs[0],), dtype=torch.float32, device="cuda" + ) + bias_, bias_storage = create_cute_tensor(bias_tensor, d_dtype, leading_dim=0) + else: + bias_tensor = None + bias_ = None + + # Resolve the compile-time activation; the device op is folded into the cubin. + if activation not in EPILOGUE_ACTIVATIONS: + raise ValueError( + f"Unsupported activation {activation!r}; " + f"expected one of {sorted(EPILOGUE_ACTIVATIONS)}" + ) + epilogue_op = EPILOGUE_ACTIVATIONS[activation]["device"] + + # Compile convolution kernel + print("Compiling kernel with cute.compile ...") + compiled_fn = compile_conv( + ncdhw, + ktrs, + input_, + filter_, + output_, + acc_dtype, + mma_tiler=mma_tiler_mn, + preferred_cluster_shape_mn=preferred_cluster_shape_mn, + fallback_cluster_shape_mn=fallback_cluster_shape_mn, + swizzle_size=swizzle_size, + raster_along=raster_along, + use_2cta_instrs=use_2cta_instrs, + upper_padding_dhw=upper_pad_dhw, + lower_padding_dhw=lower_pad_dhw, + stride_dhw=stride_dhw, + dilation_dhw=dil_dhw, + bias=bias_, + epilogue_op=epilogue_op, + ) + + # Get current CUDA stream + torch_stream = torch.cuda.Stream() + current_stream = cuda.CUstream(torch_stream.cuda_stream) + + # Run convolution. Pad/stride/dilation are passed as runtime Int32 so one + # cubin runs any pad/stride/dilation config. + print("Running Blackwell 3D convolution...") + # The inputs initialize on other streams; drain them so the kernel's + # stream never reads a tensor mid-initialization. + torch.cuda.synchronize() + compiled_fn( + input_, + filter_, + output_, + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + current_stream, + bias_, + ) + torch_stream.synchronize() + + with torch.backends.cudnn.flags(enabled=False): + if not skip_ref_check: + # Run PyTorch reference convolution + print("Running PyTorch reference 3D convolution...") + # Pytorch expects tensors to be in NCDHW format + input_ncdhw = input_tensor.permute(0, 4, 1, 2, 3) + if upper_pad_dhw != lower_pad_dhw: + # F.conv3d only supports symmetric padding, so manually pad the input + # F.pad takes padding in reverse dimension order: (W_before, W_after, H_before, H_after, D_before, D_after) + pad_arg = ( + lower_pad_dhw[2], + upper_pad_dhw[2], + lower_pad_dhw[1], + upper_pad_dhw[1], + lower_pad_dhw[0], + upper_pad_dhw[0], + ) + input_ncdhw = F.pad(input_ncdhw, pad_arg) + conv_padding = (0, 0, 0) + else: + conv_padding = upper_pad_dhw + output_ref = F.conv3d( + input_ncdhw, + filter_tensor.permute(0, 4, 1, 2, 3), + stride=stride_dhw, + padding=conv_padding, + dilation=dil_dhw, + ).to(dtype=torch.float32) + if use_bias: + output_ref = output_ref + bias_tensor.view(1, -1, 1, 1, 1) + # The kernel narrows the FP32 sum once and then activates. + output_ref = output_ref.to(dtype=torch_dtype(d_dtype)).to( + dtype=torch.float32 + ) + output_ref = EPILOGUE_ACTIVATIONS[activation]["ref"](output_ref) + # Compare results + print("Comparing results...") + + # Convert to float32 for comparison + # Transform output from (N, Z, P, Q, K) -> (N, K, Z, P, Q) + output_f32 = output_storage.permute(0, 4, 1, 2, 3).to(torch.float32) + output_ref_f32 = output_ref.to(torch.float32) + + # Verify results + torch.testing.assert_close( + output_f32, + output_ref_f32, + atol=tolerance, + rtol=1e-03, + ) + print("✓ Results match within tolerance!") + + # Benchmark if requested + if iterations > 0: + print( + f"\nBenchmarking with {warmup_iterations} warmup and {iterations} iterations..." + ) + + def generate_tensors(): + input_tensor, filter_tensor, output_tensor = prepare_tensors( + ncdhw, ktrs, (Z, P, Q), ab_dtype + ) + input_, input_storage = create_cute_tensor( + input_tensor, + ab_dtype, + leading_dim=4, + ) + filter_, filter_storage = create_cute_tensor( + filter_tensor, + ab_dtype, + leading_dim=4, + ) + output_, output_storage = create_cute_tensor( + output_tensor, + d_dtype, + leading_dim=4, + ) + # The workspace initializes on other streams; drain them so the + # benchmark stream never reads a tensor mid-initialization. + torch.cuda.synchronize() + return testing.JitArguments( + input_, + filter_, + output_, + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + current_stream, + bias_, + ) + + workspace_count = 1 + if use_cold_l2: + one_workspace_bytes = ( + input_storage.numel() * input_storage.element_size() + + filter_storage.numel() * filter_storage.element_size() + + output_storage.numel() * output_storage.element_size() + ) + workspace_count = testing.get_workspace_count( + one_workspace_bytes, warmup_iterations, iterations + ) + + # exec_time is in microseconds + exec_time = testing.benchmark( + compiled_fn, + workspace_generator=generate_tensors, + workspace_count=workspace_count, + stream=current_stream, + warmup_iterations=warmup_iterations, + iterations=iterations, + use_cuda_graphs=True, + ) + runtime_s = exec_time / 1.0e6 + fmas = (N * Z * P * Q) * K * (C * T * R * S) + flop = 2 * fmas + gflop = flop / 1.0e9 + gflops = gflop / runtime_s + + print("Average Runtime : ", exec_time / 1000, "ms") + print("GFLOPS : ", gflops) + + return exec_time + + +def _parse_comma_separated_ints(s: str) -> Tuple[int, ...]: + try: + return tuple(int(x.strip()) for x in s.split(",")) + except ValueError: + raise argparse.ArgumentTypeError( + "Invalid format. Expected comma-separated integers." + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser(description="Blackwell 3D convolution") + + # Convolution parameters + parser.add_argument( + "--ncdhw", + type=_parse_comma_separated_ints, + default=(1, 128, 32, 32, 32), + help="Input tensor shape (N,C,D,H,W)", + ) + parser.add_argument( + "--ktrs", + type=_parse_comma_separated_ints, + default=(256, 3, 3, 3), + help="Filter tensor shape components (K,T,R,S)", + ) + parser.add_argument( + "--stride_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Stride (Sd,Sh,Sw)", + ) + parser.add_argument( + "--upper_pad_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Upper padding (PadD,PadH,PadW)", + ) + parser.add_argument( + "--lower_pad_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Lower padding (PadD,PadH,PadW)", + ) + parser.add_argument( + "--dil_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Dilation (DilD,DilH,DilW)", + ) + + # Data type parameters + parser.add_argument( + "--ab_dtype", + type=cutlass.dtype, + choices=[ + cutlass.TFloat32, + cutlass.Float16, + cutlass.BFloat16, + cutlass.Int8, + cutlass.Uint8, + cutlass.Float8E4M3FN, + cutlass.Float8E5M2, + ], + default=cutlass.Float16, + help="Data type for A/B input tensors", + ) + parser.add_argument( + "--d_dtype", + type=cutlass.dtype, + choices=[ + cutlass.Float32, + cutlass.Int32, + cutlass.Float16, + cutlass.BFloat16, + cutlass.Int8, + cutlass.Uint8, + cutlass.Float8E4M3FN, + cutlass.Float8E5M2, + ], + default=cutlass.Float16, + help="Data type for output tensor D", + ) + parser.add_argument( + "--acc_dtype", + type=cutlass.dtype, + choices=[cutlass.Float32, cutlass.Float16, cutlass.Int32], + default=cutlass.Float32, + help="Accumulator data type", + ) + + # Kernel parameters + parser.add_argument( + "--mma_tiler_mn", + type=_parse_comma_separated_ints, + default=(128, 128), + help="MMA tiler shape (M,N)", + ) + parser.add_argument( + "--preferred_cluster_shape_mn", + type=_parse_comma_separated_ints, + default=(2, 1), + help="Preferred cluster shape (M,N)", + ) + parser.add_argument( + "--fallback_cluster_shape_mn", + type=_parse_comma_separated_ints, + default=(2, 1), + help="Fallback cluster shape (M,N)", + ) + parser.add_argument( + "--use_2cta_instrs", + action="store_true", + help="Enable 2CTA MMA instructions", + ) + parser.add_argument( + "--swizzle_size", + type=int, + default=1, + help="Swizzling size in the unit of cluster for improving L2 cache hit rate", + ) + parser.add_argument( + "--raster_order", + type=str, + choices=["m", "n"], + default="m", + help="Rasterization order of clusters. Only used when swizzle_size > 1", + ) + # Testing parameters + parser.add_argument( + "--tolerance", + type=float, + default=1e-02, + help="Tolerance for result comparison", + ) + parser.add_argument( + "--warmup_iterations", + type=int, + default=0, + help="Number of warmup iterations", + ) + parser.add_argument( + "--iterations", + type=int, + default=0, + help="Number of benchmark iterations", + ) + parser.add_argument( + "--skip_ref_check", + action="store_true", + help="Skip reference checking", + ) + parser.add_argument( + "--use_cold_l2", + action="store_true", + default=False, + help="Use circular buffer tensor sets to ensure L2 cold cache", + ) + parser.add_argument( + "--use_bias", + action="store_true", + help="Add a per-output-channel bias, in the output dtype, as " + "D = activation(acc + bias). Bias and no-bias are separate cubins.", + ) + parser.add_argument( + "--activation", + type=str, + default="identity", + choices=sorted(EPILOGUE_ACTIVATIONS), + help="Compile-time epilogue activation applied as D = activation(acc + " + "bias). One activation per cubin.", + ) + + args = parser.parse_args() + + run( + args.ncdhw, + args.ktrs, + args.stride_dhw, + args.upper_pad_dhw, + args.lower_pad_dhw, + args.dil_dhw, + args.ab_dtype, + args.d_dtype, + args.acc_dtype, + args.mma_tiler_mn, + args.preferred_cluster_shape_mn, + args.fallback_cluster_shape_mn, + args.swizzle_size, + args.raster_order, + args.use_2cta_instrs, + args.tolerance, + args.warmup_iterations, + args.iterations, + args.use_cold_l2, + args.skip_ref_check, + args.use_bias, + args.activation, + ) + print("PASS") diff --git a/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/blockscaled_gemm/blockscaled_gemm_dispatch.py b/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/blockscaled_gemm/blockscaled_gemm_dispatch.py index 4aaf0adfd1..2d195808da 100644 --- a/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/blockscaled_gemm/blockscaled_gemm_dispatch.py +++ b/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/blockscaled_gemm/blockscaled_gemm_dispatch.py @@ -61,6 +61,24 @@ FP4_SHIFT_BITS = 2 _FP8_DTYPES = (cutlass.Float8E4M3FN, cutlass.Float8E5M2) +_MXFP8_K64_SF_K128_TILE_SHAPE = (128, 256, 64) + + +def is_mxfp8_k64_sf_k128_tile( + tile_shape_mnk, + a_dtype, + b_dtype, + sf_dtype, + sf_vec_size, +): + """Return true when two K64 MXFP8 iterations share one K128 SF chunk.""" + return ( + tuple(tile_shape_mnk) == _MXFP8_K64_SF_K128_TILE_SHAPE + and a_dtype == b_dtype + and a_dtype in _FP8_DTYPES + and sf_dtype is cutlass.Float8E8M0FNU + and sf_vec_size == 32 + ) def make_ldmatrix_atom(operand_dtype, transpose, num_matrices=4, mixed_mode=False): @@ -192,22 +210,21 @@ def validate_blockscaled_args(args, fp4_allowed_tiles, fp8_allowed_tiles): are explicitly rejected with named diagnostics. Tile-K constraints come from the BlockScaled SF SMEM layout - (`sm120_make_smem_layout_sfa`), which requires - ``tile_K >= sf_vec_size * blk_sf == sf_vec_size * 4``: - * sf_vec_size=16 (NVFP4): tile_K must be a multiple of 64 - * sf_vec_size=32 (MXFP4 / MXFP8 / mixed): tile_K must be a multiple of 128 - A K=64 SF block cannot be filled at sf_vec_size=32 (only 2 SFs along K - fit in the K=128-required basic chunk), so tile_K=64 is rejected for - sf_vec_size=32 even though the FP4 same-dtype path otherwise allows it. + (`sm120_make_smem_layout_sfa`): + * sf_vec_size=16 (NVFP4): tile_K must be a multiple of 64 (one SF basic chunk). + * sf_vec_size=32 (MXFP4 / MXFP8 / mixed): the default SF tile_K remains + one basic chunk (=128 K-elements). The cooperative MXFP8 (128,256,64) + tile is a narrow workaround where two K64 A/B/MMA iterations share one + K128 SF chunk and each iteration consumes its matching half. + Per-path tile lists in `*_allowed_tiles` decide which tile_K values are + actually validated for a given dtype combination. """ tile = tuple(args.tile_shape_mnk) a_dtype = args.a_dtype b_dtype = args.b_dtype # Generic sf_vec_size sanity check applies to every dtype branch below. if args.sf_vec_size not in (16, 32): - raise ValueError( - f"--sf_vec_size must be 16 or 32, got {args.sf_vec_size}" - ) + raise ValueError(f"--sf_vec_size must be 16 or 32, got {args.sf_vec_size}") # Mixed-precision A/B: only the four FP4 x FP8 pairs are allowed. if a_dtype != b_dtype: if (a_dtype, b_dtype) not in MXF8F6F4_SUPPORTED_PAIRS: @@ -232,10 +249,11 @@ def validate_blockscaled_args(args, fp4_allowed_tiles, fp8_allowed_tiles): f"FP4 x FP8 mixed-precision requires --sf_dtype Float8E8M0FNU, " f"got --sf_dtype {args.sf_dtype}" ) - if tile not in fp8_allowed_tiles: + mixed_allowed_tiles = set(fp8_allowed_tiles) - {_MXFP8_K64_SF_K128_TILE_SHAPE} + if tile not in mixed_allowed_tiles: raise ValueError( f"tile_shape {tile} is not supported for FP4 x FP8 mixed-precision. " - f"Allowed mixed tile shapes: {sorted(fp8_allowed_tiles)}." + f"Allowed mixed tile shapes: {sorted(mixed_allowed_tiles)}." ) return # Same-dtype paths. @@ -271,10 +289,9 @@ def validate_blockscaled_args(args, fp4_allowed_tiles, fp8_allowed_tiles): f"tile_shape {tile} is not supported for FP4 path. " f"Allowed FP4 tile shapes: {sorted(fp4_allowed_tiles)}." ) - if args.sf_vec_size == 32 and tile[2] % 128 != 0: + if args.sf_vec_size == 32 and tile[2] != 128: raise ValueError( - f"FP4 + sf_vec_size=32 (MXFP4) requires tile_K to be a " - f"multiple of 128, " + f"FP4 + sf_vec_size=32 (MXFP4) requires tile_K=128, " f"got tile_K={tile[2]}." ) else: diff --git a/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/conv/dense_blockscaled_implicit_gemm_fprop.py b/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/conv/dense_blockscaled_implicit_gemm_fprop.py new file mode 100644 index 0000000000..90e22d3275 --- /dev/null +++ b/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/conv/dense_blockscaled_implicit_gemm_fprop.py @@ -0,0 +1,4031 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause + +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: + +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. + +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. + +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. + +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import argparse +from typing import Tuple, Type, Optional +import os +import sys + +import cuda.bindings.driver as cuda + +import torch +import torch.nn.functional as F + +import cutlass +import cutlass.cute as cute +from cutlass.cute.nvgpu import cpasync +from cutlass import testing +import cutlass.pipeline as pipeline +from cutlass.cute.runtime import from_dlpack +import cutlass.torch as cutlass_torch +import cutlass.utils.hopper_helpers as sm90_utils +import cutlass.utils.blockscaled_layout as blockscaled_utils +import cutlass.utils.blackwell_helpers as sm120_utils +from pathlib import Path + +if __name__ == "__main__": + # `helpers` sits at the examples/CuTeDSL root; running this file as a + # script only puts its own directory on sys.path. + cutedsl_dir = str(Path(__file__).resolve().parents[4]) + if cutedsl_dir not in sys.path: + sys.path.insert(0, cutedsl_dir) + +from helpers.dynamic_persistent_tile_scheduler import ( + ClcDynamicPersistentTileScheduler, + ClcDynamicPersistentTileSchedulerParams, +) +from cutlass.memory import SmemAllocator +from cutlass.tensor_utils import LayoutEnum + +if __name__ == "__main__": + current_dir = os.path.dirname(os.path.abspath(__file__)) + sys.path.insert(0, os.path.join(current_dir, "../../../")) + +# SM120 block-scaled GEMM dispatch helpers (sibling utility module). Try the +# namespace-package path first, which resolves when the examples root is on +# sys.path; fall back to the bare local import, which resolves when only this +# file's own directory is on sys.path[0]. +try: + from blackwell_geforce.kernel.blockscaled_gemm.blockscaled_gemm_dispatch import ( + FP4_SHIFT_BITS, + make_ldmatrix_atom, + make_sm120_blockscaled_mma_op, + ) +except ImportError: + from blockscaled_gemm_dispatch import ( # noqa: E402 + FP4_SHIFT_BITS, + make_ldmatrix_atom, + make_sm120_blockscaled_mma_op, + ) + +# Conv helpers shared with the legacy SM120 conv, imported on the same two paths. +try: + from blackwell_geforce.kernel.conv.dense_implicit_gemm_fprop import ( + Sm120PersistentDenseImplicitGemmFpropKernel, + _check_im2col_descriptor_limits, + _check_tensor_alignment, + _compute_im2col_params, + create_cute_tensor, + _parse_comma_separated_ints, + _rt_conv_scalars, + compute_zpq, + prepare_tensors, + ) +except ImportError: + from dense_implicit_gemm_fprop import ( # noqa: E402 + Sm120PersistentDenseImplicitGemmFpropKernel, + _check_im2col_descriptor_limits, + _check_tensor_alignment, + _compute_im2col_params, + create_cute_tensor, + _parse_comma_separated_ints, + _rt_conv_scalars, + compute_zpq, + prepare_tensors, + ) + +""" +SM120 NVFP4 fprop in CuTe DSL. + +The device mainloop is intentionally kept aligned with +dense_blockscaled_gemm_persistent_pingpong.py. Convolution is represented as an +implicit GEMM view: +- A: NDHWC im2col -> (N*Z*P*Q, C*T*R*S, 1) +- B: KTRSC -> (K, C*S*R*T, 1) +- D: NZPQK im2col store -> (N*Z*P*Q, K, 1) + +SFB follows the SM120 block-scaled GEMM path. SFA is loaded by a single-CTA +async-copy producer into the same SMEM layout that the SM120 GEMM math path +consumes. Conv geometry (T/R/S, pad/stride/dilation, N/D/H/W) is passed as +runtime scalars, the channel count among them: SFB's per-position span is derived +from it at runtime, so one compiled kernel serves every C. +""" + + +# ///////////////////////////////////////////////////////////////////////////// +# Helpers to parse args +# ///////////////////////////////////////////////////////////////////////////// + + +# ///////////////////////////////////////////////////////////////////////////// +# Host setup and device kernel launch +# ///////////////////////////////////////////////////////////////////////////// + + +class Sm120BlockScaledPersistentDenseImplicitGemmFpropKernel( + Sm120PersistentDenseImplicitGemmFpropKernel +): + def __init__( + self, + acc_dtype, + sf_vec_size, + tile_shape_mnk, + ): + super().__init__(acc_dtype, tile_shape_mnk) + self.sf_vec_size = sf_vec_size + + # Override the warp specialization: SFA cannot be loaded by a TMA (a K tile + # reads a different pixel per filter position, so the M address is not a + # regular box), so it gets its own 4-warp cp.async group right after the MMA + # warps. The TMA load, CLC scheduler, and residual G2S warps then share the + # round-up warpgroup; the residual warp reuses an otherwise-idle slot, so + # enabling residual grows neither the CTA nor the register budget. + self.sfa_cpasync_warp_group_start = self.num_mma_warps + self.sfa_cpasync_warp_group_size = 4 + self.tma_load_warp_id = ( + self.sfa_cpasync_warp_group_start + self.sfa_cpasync_warp_group_size + ) + self.sched_warp_id = self.tma_load_warp_id + 1 + self.residual_tma_warp_id = self.sched_warp_id + 1 + self.threads_per_cta = ( + self.num_mma_warps + + self.sfa_cpasync_warp_group_size # SFA cp.async DMA warp group + + 1 # dedicated TMA load warp + + self.num_sched_warps # CLC scheduler warp + ) * self.num_threads_per_warp + + # epi_tile is chosen in _setup_attributes once has_residual is known: a + # residual epilogue needs a separate sResidual staging buffer, so it uses a + # subtiled epi_tile to keep the epilogue smem (sD + sResidual) small enough + # to preserve the A/B pipeline depth; without a residual the full (128,128) + # epi_tile is fastest (single subtile, deepest A/B pipeline). + self.epi_tile = None + + # The CLC scheduler warp plus the round-up unused warps form a whole extra + # warpgroup whose only real work is the lightweight CLC producer loop; give + # it the register-realloc floor so the freed budget goes to the MMA + # warpgroups. Its value is shared by the scheduler and unused branches since + # they belong to the same warpgroup. + self.sched_register_requirement = 24 + self.mma_register_requirement = 224 + + def _setup_attributes(self): + # Pick the epilogue tile from whether a residual is fused. Without a + # residual the full-CTA epi_tile stores D in one subtile and leaves the + # most smem for the A/B pipeline. With a residual, sResidual doubles the + # epilogue smem, so a subtiled epi_tile keeps sD + sResidual small enough + # to hold the A/B pipeline depth; (64,64) measured fastest across the + # residual shapes. + epi_m, epi_n = (64, 64) if self.has_residual else (128, 128) + # The epilogue subtile cannot exceed the CTA tile, and its N extent has to + # divide tile_n. The storing warpgroup takes its subtile count from the + # tiled gmem tensor's rest mode, which rounds up, while the accumulator + # index it builds per subtile is epi_n * (epi_tile_n // mma_tile_n) + ... -- + # so a remainder makes that index run past the accumulator fragment. The + # residual skip count and the epilogue stage budget divide instead, which + # rounds down. Halving epi_n until it divides keeps all three in agreement + # for every tile_n. It only fires on the residual branch at tile_n=96, where + # the starting 64 does not divide and halves to 32; without a residual the + # starting 128 clamps to 96, which divides itself. + epi_m = min(epi_m, self.tile_shape_mnk[0]) + epi_n = min(epi_n, self.tile_shape_mnk[1]) + while self.tile_shape_mnk[1] % epi_n != 0: + epi_n //= 2 + self.epi_tile = (epi_m, epi_n) + # The SF SMEM atom (BlockScaledBasicChunk) is indivisible along N: one + # chunk spans blk_mn=128 output channels, so a tile_n below 128 still + # stages a whole chunk and n_tiles_per_sf_n_tile consecutive N tiles share + # it; the consumer offsets into its own slice of the chunk. + self.sf_tile_n = ((self.tile_shape_mnk[1] + 127) // 128) * 128 + self.n_tiles_per_sf_n_tile = self.sf_tile_n // self.tile_shape_mnk[1] + + mma_op, use_mxf8f6f4 = make_sm120_blockscaled_mma_op( + self.a_dtype, + self.b_dtype, + self.acc_dtype, + self.sf_dtype, + self.sf_vec_size, + ) + # mixed_mode carries the full mixed FP4 x FP8 machinery (Int8 SMEM/TMA + # recast + FP4_SHIFT_BITS unpack) for a future mixed-dtype path. run_conv + # currently gates ab_dtype to Float4E2M1FN, so a_dtype == b_dtype and + # mixed_mode is always False here -- the shift branches are unreachable in + # this example but kept so the mixed path can be enabled without a rewrite. + self.mixed_mode = self.a_dtype != self.b_dtype + # a_fp4_in_mixed / b_fp4_in_mixed: this side carries an FP4 operand in the + # mixed FP4 x FP8 mode, so SMEM/TMA see Int8 storage and the mma.sync + # consumer needs the LDSM b4x16_p64 unpack + register `<< FP4_SHIFT_BITS`. + a_fp4_in_mixed = self.mixed_mode and self.a_dtype.width < 8 + b_fp4_in_mixed = self.mixed_mode and self.b_dtype.width < 8 + self.smem_alloc_a_dtype = cutlass.Int8 if a_fp4_in_mixed else self.a_dtype + self.smem_alloc_b_dtype = cutlass.Int8 if b_fp4_in_mixed else self.b_dtype + # `internal_type` for `_make_tma_atoms_and_tensors` is None when the dtype + # already matches (TMA sees the native dtype), Int8 when we recast for FP4. + self.tma_internal_a_dtype = cutlass.Int8 if a_fp4_in_mixed else None + self.tma_internal_b_dtype = cutlass.Int8 if b_fp4_in_mixed else None + atom_shape = (2, 2, 1) + atom_layout = cute.make_layout(atom_shape) + permutation_mnk = sm120_utils.get_permutation_mnk( + self.tile_shape_mnk, self.sf_vec_size, use_mxf8f6f4 + ) + self.tiled_mma = cute.make_tiled_mma( + mma_op, + atom_layout, + permutation_mnk=permutation_mnk, + ) + + # Scale-factor tiles decouple from the A/B K tile for MXFP8. The SF SMEM + # atom (BlockScaledBasicChunk) spans 4 SF blocks along K, so at vec32 one + # atom covers 4*32 = 128 channels: the SF K tile must be 128 even when the + # A/B K tile is 64. One K128 SF chunk then serves two K64 A/B tiles, and + # the consumer offsets into the live half via sf_base. NVFP4 keeps + # ab_k_tiles_per_sf_k_tile == 1, so every SF path below is a no-op for it. + self.sf_tile_k = sf_k_tile_channels(self.tile_shape_mnk[2], self.sf_vec_size) + self.ab_k_tiles_per_sf_k_tile = self.sf_tile_k // self.tile_shape_mnk[2] + # The SF tile takes the rounded-up N from above and the K cadence from here, + # so a narrow tile_n still stages a whole chunk and MXFP8 still stages one + # whole chunk along K. + self.sf_tile_shape_mnk = ( + self.tile_shape_mnk[0], + self.sf_tile_n, + self.sf_tile_k, + ) + + self.cta_layout_mnk = cute.make_layout(self.cluster_shape_mnk) + + # Compute stage before compute smem layout + self.ab_stage, self.epi_stage = self._compute_stages( + self.tile_shape_mnk, + self.smem_alloc_a_dtype, + self.smem_alloc_b_dtype, + self.sf_dtype, + self.epi_tile, + self.d_dtype, + self.smem_capacity, + self.occupancy, + self.has_residual, + ) + + # Residual G2S ring depth, independent of the D-store epi_stage. A single + # buffer already hides the residual load: the dedicated residual warp runs + # its own acquire/commit loop and stays a subtile ahead of the epilogue + # consumer, which releases each stage right after the S2R read (before the + # add), so the buffer turns over fast. A deeper ring only steals smem from + # the A/B pipeline and lowers ab_stage, so keep it at 1. + self.res_stage = 1 + + assert self.epi_stage > 0, ( + "epi_stage <= 0, no enough shared memory. This case will be skipped." + ) + + ( + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.sfa_smem_layout_staged, + self.sfb_smem_layout_staged, + self.epi_smem_layout_staged, + self.res_smem_layout_staged, + ) = self._make_smem_layouts( + self.tile_shape_mnk, + self.epi_tile, + self.smem_alloc_a_dtype, + self.a_layout, + self.smem_alloc_b_dtype, + self.b_layout, + self.ab_stage, + self.d_dtype, + self.d_layout, + self.epi_stage, + self.res_stage, + self.sf_vec_size, + self.tiled_mma, + self.sf_tile_shape_mnk, + ) + # MMA-side view of the staged SFB chunk. The TMA stages whole 128-wide + # chunks, but the mma only consumes the live tile_n columns, so the + # consumer reads the same SMEM through a tile_n-wide layout whose modes + # then match the B fragment's. Only N narrows: K stays on the SF cadence, + # which is one whole chunk and can be wider than the A/B K tile. + self.sfb_mma_layout_staged = _sm120_make_smem_layout_sfb( + self.tiled_mma, + (self.tile_shape_mnk[0], self.tile_shape_mnk[1], self.sf_tile_k), + self.sf_vec_size, + self.ab_stage, + ) + + @cute.jit + def __call__( + self, + a: cute.Tensor, + b: cute.Tensor, + sfa: cute.Tensor, + sfb: cute.Tensor, + d: cute.Tensor, + bias: Optional[cute.Tensor], + residual: Optional[cute.Tensor], + alpha: cutlass.Float32, + beta: cutlass.Constexpr, + rt_upper_pad_d: cutlass.Int32, + rt_upper_pad_h: cutlass.Int32, + rt_upper_pad_w: cutlass.Int32, + rt_lower_pad_d: cutlass.Int32, + rt_lower_pad_h: cutlass.Int32, + rt_lower_pad_w: cutlass.Int32, + rt_stride_d: cutlass.Int32, + rt_stride_h: cutlass.Int32, + rt_stride_w: cutlass.Int32, + rt_dil_d: cutlass.Int32, + rt_dil_h: cutlass.Int32, + rt_dil_w: cutlass.Int32, + rt_filter_t: cutlass.Int32, + rt_filter_r: cutlass.Int32, + rt_filter_s: cutlass.Int32, + max_active_clusters: cutlass.Constexpr, + stream: cuda.CUstream, + epilogue_op: cutlass.Constexpr = lambda x: x, + ): + """Execute fprop using the SM120 block-scaled GEMM mainloop. + + alpha is an FP32 runtime scalar applied to the accumulator before bias + as D = act(alpha * acc + bias + beta * residual). It is passed at both + cute.compile and the runtime launch, so its value stays out of the + mangled name and one cubin serves any alpha. + + beta is a compile-time Constexpr, so beta == 0 folds the entire residual + path (dedicated load warp, smem staging, TMA, add) away via DCE and + beta != 0 compiles it in. A separate cubin is emitted per beta value; the + residual tensor shares the output's shape/layout/dtype. + + Pad/stride/dilation arrive as boxed cutlass.Int32 scalars so they lower + to runtime SSA: one compiled cubin serves any pad/stride/dilation + config. The SAME values feed BOTH the im2col A descriptor corner (host, + _compute_im2col_params) and the device SFA cp.async coord reconstruction + (sfa_cpasync_copy_tile); runtime-ing only one side would desync A's load + coords from SFA's, zeroing the accumulator on any non-default config. + """ + + # setup static attributes before smem/grid/tma computation + self.a_dtype = a.element_type + self.b_dtype = b.element_type + self.d_dtype = d.element_type + self.sf_dtype = sfa.element_type + + # Bias opt-in: caller supplies a length-K (output-channel) tensor. has_bias + # is a compile-time branch, so no-bias compiles the bias path away. + self.has_bias: bool = bias is not None + # The bias register fragment is allocated with the bias tensor's dtype and + # upconverted to FP32 for the add. Require bias to match the output dtype. + if cutlass.const_expr(self.has_bias and bias.element_type is not self.d_dtype): + raise RuntimeError( + f"bias dtype ({bias.element_type}) must match output " + f"dtype ({self.d_dtype})" + ) + + # Residual opt-in: gated by the compile-time beta. beta == 0 compiles the + # whole residual path away; beta != 0 requires a residual tensor sharing + # the output's shape/layout/dtype (it is TMA-loaded through the output's + # own store descriptor, run in reverse as a G2S load). + self.beta = beta + self.has_residual: bool = beta != 0.0 + if cutlass.const_expr(self.has_residual and residual is None): + raise RuntimeError("beta != 0 requires a residual tensor") + if cutlass.const_expr( + self.has_residual and residual.element_type is not self.d_dtype + ): + raise RuntimeError( + f"residual dtype ({residual.element_type}) must match output " + f"dtype ({self.d_dtype})" + ) + + if cutlass.const_expr(a.leading_dim != 4): + raise RuntimeError("Input must be NDHWC with C contiguous") + if cutlass.const_expr(b.leading_dim != 4): + raise RuntimeError("Filter must be KTRSC with C contiguous") + if cutlass.const_expr(d.leading_dim != 4): + raise RuntimeError("Output must be NZPQK with K contiguous") + + # In the flattened GEMM view all operands are K-contiguous / row-major. + self.a_layout = LayoutEnum.ROW_MAJOR + self.b_layout = LayoutEnum.ROW_MAJOR + self.d_layout = LayoutEnum.ROW_MAJOR + + self._setup_attributes() + + def add_dummy_batch_dimension(tensor): + return cute.make_tensor( + tensor.iterator, cute.append(tensor.layout, cute.make_layout(1)) + ) + + # Filter T/R/S arrive as boxed cutlass.Int32 (not b.shape, which trace- + # time resolves to compile-time ints). Runtime extents keep the B im2col + # K layout at full rank 4 (c,s,r,t): when an extent is a compile-time 1 + # the tile partition collapses that mode, but a runtime extent cannot be + # proven == 1 at compile time so the mode is retained — matching A's + # always-rank-4 K traversal and letting one cubin serve any T/R/S. + rt_filter_trs = (rt_filter_t, rt_filter_r, rt_filter_s) + rt_upper_pad_dhw = (rt_upper_pad_d, rt_upper_pad_h, rt_upper_pad_w) + rt_lower_pad_dhw = (rt_lower_pad_d, rt_lower_pad_h, rt_lower_pad_w) + rt_stride_dhw = (rt_stride_d, rt_stride_h, rt_stride_w) + rt_dil_dhw = (rt_dil_d, rt_dil_h, rt_dil_w) + + ( + lower_corner_whd, + upper_corner_whd, + lower_padding_whd, + upper_padding_whd, + stride_whd, + lower_srt, + stride_srt, + ) = _compute_im2col_params( + rt_filter_trs, + rt_upper_pad_dhw, + rt_lower_pad_dhw, + rt_stride_dhw, + rt_dil_dhw, + ) + + # A im2col: (N,D,H,W,C) -> ((W,H,D,N), C), descriptor K=(C,S,R,T). + a_gemm = cute.make_tensor( + a.iterator, + cute.select(a.layout, mode=[3, 2, 1, 0, 4]), + ) + a_gemm = cute.group_modes(a_gemm, begin=0, end=4) + + # B: (K,T,R,S,C) -> reorder to (K, (C,S,R,T)) for K-coordinate indexing. + # Reorder the mark_layout_dynamic filter tensor's modes instead of + # rebuilding a compact layout from scalars: the reordered view inherits + # the tensor's runtime C/T/R/S extents and strides, so no channel or + # filter extent is baked into the cubin. + b_gemm = cute.make_tensor( + b.iterator, cute.select(b.layout, mode=[0, 4, 3, 2, 1]) + ) + b_gemm = cute.group_modes(b_gemm, begin=1, end=5) + + # D im2col store: (N,Z,P,Q,K) -> ((Q,P,Z,N), K). + d_gemm = cute.make_tensor( + d.iterator, cute.select(d.layout, mode=[3, 2, 1, 0, 4]) + ) + d_gemm = cute.group_modes(d_gemm, begin=0, end=4) + + # Build mBias tensor: a per-output-channel bias broadcast to the same + # ((Q,P,Z,N), K, 1) = (M, N, L) profile as the output. The bias varies + # along the output channel K (GEMM-N, real stride 1) and broadcasts across + # every spatial output position (GEMM-M = Q,P,Z,N modes carry stride 0), so + # the epilogue reads it through the same partition_C chain as the + # accumulator and each thread lands on its own N-column bias scalar. The M + # (spatial) modes are stride-0 broadcast, so a static M extent (the CTA + # tile size) keeps the partitioned smem-read layout static for a vectorized + # load while the N axis carries a runtime stride so one cubin serves any + # output-channel count. + if cutlass.const_expr(self.has_bias): + mBias_layout = cute.make_layout( + ( + (self.tile_shape_mnk[0], 1, 1, 1), + cute.size(d, mode=[4]), + 1, + ), + stride=((0, 0, 0, 0), 1, 0), + ) + mBias_mnl = cute.make_tensor(bias.iterator, mBias_layout) + else: + mBias_mnl = None + + # SFA is an unswizzled input-activation scale tensor in + # (N*D*H*W, ceil(C / sf_vec_size), 1). The async-copy producer maps each + # output tile and TRS coordinate back to this input-space tensor before + # writing the SM120 SFA SMEM layout consumed by the MMA math path. + sfa_tensor = sfa + + # Setup SFB tensor by filling B tensor to scale factor atom layout. + # ((Atom_N, Rest_N),(Atom_K, Rest_K),RestL) + # + # The channel extent is the span SFB was allocated at, not C: the host pads + # each filter position out to a whole number of SF K tiles. That padding is + # what keeps the flat K mode usable -- it puts every filter position on a + # tile boundary, so a K tile cut from the mode stays inside one position + # while A and B take zero-filled channels for the tail, and the consumer + # addresses a tile by its flat index. The span is derived from the runtime + # channel count, so it never reaches the compiled code. + sfb_c_span = cute.ceil_div(b.shape[4], self.sf_tile_k) * self.sf_tile_k + sfb_shape = ( + b.shape[0], + sfb_c_span * b.shape[3] * b.shape[2] * b.shape[1], + 1, + ) + sfb_layout = blockscaled_utils.tile_atom_to_shape_SF( + sfb_shape, self.sf_vec_size + ) + sfb_tensor = cute.make_tensor(sfb.iterator, sfb_layout) + + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, 0)) + tma_atom_a, tma_tensor_a = cpasync.make_im2col_tma_atom( + cpasync.CopyBulkTensorIm2ColG2SOp(), + a_gemm, + a_smem_layout, + (self.tile_shape_mnk[0], self.tile_shape_mnk[2]), + lower_corner_whd, + upper_corner_whd, + lower_padding_whd, + upper_padding_whd, + stride_whd, + lower_srt, + stride_srt, + internal_type=self.tma_internal_a_dtype, + ) + tma_tensor_a = add_dummy_batch_dimension(tma_tensor_a) + + # B cta_tiler K mode is nested (tileK,), not flat tileK: b_gemm's K is the + # nested (C,S,R,T) im2col group with runtime S/R/T. A nested tiler tiles + # only the first (channel) level -> a static (tileK,) tile with runtime + # S/R/T left in the rest, so the TMA CTA V-map (composition of the tiler + # with identity(gmem.shape)) stays static while C and S/R/T stay runtime. + # A flat tileK keeps the whole runtime nest and fails the static V-map + # check. + tma_atom_b, tma_tensor_b = self._make_tma_atoms_and_tensors( + b_gemm, + self.b_smem_layout_staged, + (self.tile_shape_mnk[1], (self.tile_shape_mnk[2],)), + 1, + internal_type=self.tma_internal_b_dtype, + ) + tma_tensor_b = add_dummy_batch_dimension(tma_tensor_b) + + tma_atom_sfb, tma_tensor_sfb = self._make_tma_atoms_and_tensors( + sfb_tensor, + self.sfb_smem_layout_staged, + (self.sf_tile_shape_mnk[1], self.sf_tile_shape_mnk[2]), + 1, + internal_type=cutlass.Int16, + ) + + epi_smem_layout = cute.slice_(self.epi_smem_layout_staged, (None, None, 0)) + tma_atom_d, tma_tensor_d = cpasync.make_im2col_tma_atom( + cpasync.CopyBulkTensorIm2ColS2GOp(), + d_gemm, + epi_smem_layout, + self.epi_tile, + ) + tma_tensor_d = cute.coalesce(tma_tensor_d, target_profile=(1, 1)) + tma_tensor_d = add_dummy_batch_dimension(tma_tensor_d) + + # Residual im2col G2S load: the exact inverse of the output's im2col S2G + # store, so it builds the same ((Q,P,Z,N), K) GEMM view (from the residual + # tensor's own iterator) and reuses the epilogue smem layout. Load mode + # needs the full im2col corner set; the residual is a 1x1x1 identity tap + # (same shape/layout/dtype as the output) so every corner/pad is 0, DHW + # stride is 1, and the SRT lower/stride are 0/1. + if cutlass.const_expr(self.has_residual): + residual_gemm = cute.make_tensor( + residual.iterator, cute.select(residual.layout, mode=[3, 2, 1, 0, 4]) + ) + residual_gemm = cute.group_modes(residual_gemm, begin=0, end=4) + tma_atom_residual, tma_tensor_residual = cpasync.make_im2col_tma_atom( + cpasync.CopyBulkTensorIm2ColG2SOp(), + residual_gemm, + epi_smem_layout, + self.epi_tile, + lower_corner_whd=(0, 0, 0), + upper_corner_whd=(0, 0, 0), + lower_padding_whd=(0, 0, 0), + upper_padding_whd=(0, 0, 0), + stride_whd=(1, 1, 1), + lower_srt=(0, 0, 0), + stride_srt=(1, 1, 1), + ) + tma_tensor_residual = cute.coalesce( + tma_tensor_residual, target_profile=(1, 1) + ) + tma_tensor_residual = add_dummy_batch_dimension(tma_tensor_residual) + else: + tma_atom_residual, tma_tensor_residual = None, None + + tile_sched_params, grid = self._compute_grid( + tma_tensor_d, + self.tile_shape_mnk, + max_active_clusters, + ) + + @cute.struct + class SharedStorage: + mainloop_pipeline_array_ptr: cute.struct.MemRange[ + cutlass.Int64, self.ab_stage * 2 + ] + # Two math warpgroups take turns at two points each -- before the + # mainloop and before the epilogue -- so the array holds 2 * 2 mbarriers. + math_wg_order_barrier_array_ptr: cute.struct.MemRange[cutlass.Int64, 4] + clc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_clc_stage * 2] + clc_response: cute.struct.Align[ + cute.struct.MemRange[cutlass.Int32, self.num_clc_stage * 4], 16 + ] + sA: cute.struct.Align[ + cute.struct.MemRange[ + self.smem_alloc_a_dtype, cute.cosize(self.a_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + sB: cute.struct.Align[ + cute.struct.MemRange[ + self.smem_alloc_b_dtype, cute.cosize(self.b_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + sSFA: cute.struct.Align[ + cute.struct.MemRange[ + self.sf_dtype, cute.cosize(self.sfa_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + sSFB: cute.struct.Align[ + cute.struct.MemRange[ + self.sf_dtype, cute.cosize(self.sfb_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + sD: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, cute.cosize(self.epi_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + # Ring of cta_tile_n bias rows: the TMA load warp cp.async's each + # tile's row into the stage its producer state points at, and the + # tile's owning math warpgroup reads it back through its consumer + # state. Zero length when no bias is supplied so it costs no smem. + sBias: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, + self.bias_stage * self.tile_shape_mnk[1] if self.has_bias else 0, + ], + self.buffer_align_bytes, + ] + # Residual staging: same per-subtile atom as sD, single-buffered + # (res_stage=1) so it costs minimal smem and leaves the A/B pipeline + # deep. The S2R read lines up with the acc fragment per subtile. Zero + # length (no smem) when beta == 0. + sResidual: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, + cute.cosize(self.res_smem_layout_staged) + if self.has_residual + else 0, + ], + self.buffer_align_bytes, + ] + # PipelineTmaAsync mbarriers for the residual G2S load (full+empty per + # stage). Sized by res_stage. Zero when beta == 0. + residual_pipeline_array_ptr: cute.struct.MemRange[ + cutlass.Int64, self.res_stage * 2 if self.has_residual else 0 + ] + # PipelineCpAsync mbarriers for the bias staging ring (full+empty per + # stage). Zero when no bias is supplied. + bias_pipeline_array_ptr: cute.struct.MemRange[ + cutlass.Int64, self.bias_stage * 2 if self.has_bias else 0 + ] + + self.shared_storage = SharedStorage + + self.threads_per_cta = (self.threads_per_cta + 127) // 128 * 128 + + self.kernel( + tma_atom_a, + tma_tensor_a, + tma_atom_b, + tma_tensor_b, + sfa_tensor, + tma_atom_sfb, + tma_tensor_sfb, + tma_atom_d, + tma_tensor_d, + tma_atom_residual, + tma_tensor_residual, + mBias_mnl, + alpha, + beta, + self.tiled_mma, + self.cta_layout_mnk, + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.sfa_smem_layout_staged, + self.sfb_smem_layout_staged, + self.sfb_mma_layout_staged, + self.epi_smem_layout_staged, + self.res_smem_layout_staged, + tile_sched_params, + cute.size(a, mode=[0]), + cute.size(a, mode=[1]), + cute.size(a, mode=[2]), + cute.size(a, mode=[3]), + cute.size(a, mode=[4]), + cute.size(d, mode=[1]), + cute.size(d, mode=[2]), + cute.size(d, mode=[3]), + rt_stride_dhw[0], + rt_stride_dhw[1], + rt_stride_dhw[2], + rt_lower_pad_dhw[0], + rt_lower_pad_dhw[1], + rt_lower_pad_dhw[2], + rt_dil_dhw[0], + rt_dil_dhw[1], + rt_dil_dhw[2], + rt_filter_trs[0], + rt_filter_trs[1], + rt_filter_trs[2], + epilogue_op, + ).launch( + grid=grid, + block=[self.threads_per_cta, 1, 1], + cluster=[1, 1, 1], + stream=stream, + max_number_threads=[self.threads_per_cta, 1, 1], + min_blocks_per_mp=1, + ) + return + + @cute.jit + def advance(self, state: pipeline.PipelineState, iterations): + if iterations < state.stages and ((state._index + iterations) >= state.stages): + state._phase ^= 1 + if ( + iterations >= state.stages + and (((state._index + iterations) // state.stages) % 2) == 1 + ): + state._phase ^= 1 + state._index = (state._index + iterations) % state.stages + state._count += iterations + return state + + def _make_b_k_coord( + self, + c_tile_idx, + s_idx, + r_idx, + t_idx, + ): + # After local_tile splits the channel into K-tiles, B's im2col K view is + # (c_tile, (s, r, t)): the channel-tile is a top-level mode and S/R/T are + # a nested group. All four modes are always present because the B layout + # is built with runtime C and T/R/S extents — the partition cannot prove + # any extent == 1 at compile time, so no mode ever collapses, and the + # K-coord always carries (c_tile, s, r, t) as a flat 4-mode group. + return (c_tile_idx, s_idx, r_idx, t_idx) + + @cute.jit + def make_and_init_order_barrier(self, order_mbar_ptr, group_id): + StagesPerMathWarpGroup = 2 + return pipeline.PipelineOrder.create( + barrier_storage=order_mbar_ptr, + depth=StagesPerMathWarpGroup, + length=2, + group_id=group_id, + producer_group=pipeline.CooperativeGroup( + pipeline.Agent.Thread, + 128, + ), + defer_sync=True, + ) + + @cute.jit + def sfa_cpasync_copy_tile( + self, + mSFA_mkl: cute.Tensor, + sSFA: cute.Tensor, + tile_coord_mnl, + sf_stage, + sfa_c_tile_idx, + s_coord, + r_coord, + t_coord, + tidx, + input_N: cutlass.Int32, + input_C: cutlass.Int32, + input_D: cutlass.Int32, + input_H: cutlass.Int32, + input_W: cutlass.Int32, + output_Z: cutlass.Int32, + output_P: cutlass.Int32, + output_Q: cutlass.Int32, + rt_stride_d: cutlass.Int32, + rt_stride_h: cutlass.Int32, + rt_stride_w: cutlass.Int32, + rt_lower_pad_d: cutlass.Int32, + rt_lower_pad_h: cutlass.Int32, + rt_lower_pad_w: cutlass.Int32, + rt_dil_d: cutlass.Int32, + rt_dil_h: cutlass.Int32, + rt_dil_w: cutlass.Int32, + ): + sfa_atom_copy = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(), + mSFA_mkl.element_type, + num_bits_per_copy=32, + ) + # The whole 4-warp DMA group cooperates on the SFA copy: each of the 128 + # lanes owns exactly one M token of the 128 x (K/sf_vec_size) tile, so a + # lane loads sf_k_per_ktile scale factors (vec16 -> 8, vec32 -> 4) as + # sf_k_per_ktile//4 chunked 4-byte cp.async copies. + # (warps 8..11 -> local_tid 0..127 via tidx % 128). + local_tid = tidx % 128 + sfa_predicate_tensor = cute.make_rmem_tensor( + cute.make_layout((1,)), + cutlass.Boolean, + ) + + ZPQ = output_Z * output_P * output_Q + PQ = output_P * output_Q + DHW = input_D * input_H * input_W + HW = input_H * input_W + cta_m_tile_offset = tile_coord_mnl[0] * self.tile_shape_mnk[0] + # A CTA loads one full SF K tile (self.sf_tile_k channels of scale + # factors) per lane, which the SMEM BlockScaledBasicChunk atom holds. When + # the SF K tile covers several A/B K tiles (ab_k_tiles_per_sf_k_tile > 1), + # the paired A/B tiles share the same filter position, so their SF are + # contiguous in the unswizzled gmem and one SF K tile fills them all; the + # consumer offsets into the live A/B-tile half via sf_base. + sf_k_per_ktile = self.sf_tile_k // self.sf_vec_size + # SFA SMEM uses BlockScaledBasicChunks of blk_sf=4 consecutive scale + # factors (one per channel group). Within a chunk the SF for consecutive + # channel groups are contiguous (identity layout) in both the unswizzled + # global source and the SMEM atom, so a chunk is a single 4-byte + # cp.async; consecutive chunks are separated by the K basic-block stride + # in SMEM, so we issue one copy per chunk. The chunk count adapts to + # sf_vec_size (vec16 -> 8 SF -> 2 chunks; vec32 -> 4 SF -> 1 chunk). + sf_blk = 4 + sf_chunk_cnt = sf_k_per_ktile // sf_blk + sf_k_base = (sfa_c_tile_idx // self.ab_k_tiles_per_sf_k_tile) * sf_k_per_ktile + + # SMEM M coordinate is the 2-level ((32, 4)) mode; the global token of a + # lane is (inner + 32 * outer) == local_tid, matching the layout that the + # SM120 math path reads back. + sf_m_inner = local_tid % 32 + sf_m_outer = local_tid // 32 + + rel_token = local_tid + m_global = cta_m_tile_offset + rel_token + n_idx = m_global // ZPQ + zpq_rem = m_global % ZPQ + z_idx = zpq_rem // PQ + pq_rem = zpq_rem % PQ + p_idx = pq_rem // output_Q + q_idx = pq_rem % output_Q + m_valid = m_global < (input_N * ZPQ) + n_clamped = n_idx if m_valid else 0 + + d_in = z_idx * rt_stride_d - rt_lower_pad_d + t_coord * rt_dil_d + h_in = p_idx * rt_stride_h - rt_lower_pad_h + r_coord * rt_dil_h + w_in = q_idx * rt_stride_w - rt_lower_pad_w + s_coord * rt_dil_w + + d_cl = d_in if d_in >= 0 else 0 + d_cl = d_cl if d_cl < input_D else 0 + h_cl = h_in if h_in >= 0 else 0 + h_cl = h_cl if h_cl < input_H else 0 + w_cl = w_in if w_in >= 0 else 0 + w_cl = w_cl if w_cl < input_W else 0 + sfa_m_addr = n_clamped * DHW + d_cl * HW + h_cl * input_W + w_cl + + sfa_pred_val = cutlass.Boolean(0) + if m_valid: + if d_in >= 0 and d_in < input_D: + if h_in >= 0 and h_in < input_H: + if w_in >= 0 and w_in < input_W: + sfa_pred_val = cutlass.Boolean(1) + sfa_predicate_tensor[0] = sfa_pred_val + + sfa_m_slice = mSFA_mkl[(sfa_m_addr, None, 0)] + + # A pixel owns only ceil(C / sf_vec_size) real scale factors, so a K tile + # wider than C has trailing chunks with no backing storage. Reading one + # would walk into the next pixel's scale factors, and off the end of the + # buffer on the last pixel. Those chunks pair with the A channels past C, + # which the im2col TMA zero-fills, so their scale factors never reach the + # result -- but they must still name an in-bounds address, and must land a + # finite value in SMEM, because the MMA applies the scale before it can + # know the operand is zero. Re-reading the pixel's first chunk gives both. + sf_real_blocks = cute.ceil_div(input_C, self.sf_vec_size) + + for c_chunk in cutlass.range_constexpr(sf_chunk_cnt): + local_sf_k = sf_k_base + c_chunk * sf_blk + gmem_sf_k = local_sf_k if local_sf_k < sf_real_blocks else 0 + tAgSFA_slice = cute.make_tensor( + sfa_m_slice.iterator + gmem_sf_k, + layout=cute.make_layout((sf_blk,)), + ) + # Slice the whole chunk: span the K basic-block (mma_nsf) and the + # blk_sf//mma_nsf middle mode for this chunk c_chunk. Their combined + # extent is always blk_sf and is contiguous from the chunk base, so a + # flat (blk_sf,) layout over the slice iterator is the identity copy. + tAsSFA_slot = sSFA[ + ( + ( + (sf_m_inner, sf_m_outer), + 0, + ), + ((0, None), None, c_chunk), + sf_stage, + ) + ] + tAsSFA_slice = cute.make_tensor( + tAsSFA_slot.iterator, + layout=cute.make_layout((sf_blk,)), + ) + + cute.copy_atom_call( + sfa_atom_copy, + tAgSFA_slice, + tAsSFA_slice, + pred=sfa_predicate_tensor, + ) + + @cute.jit + def sfa_cpasync_copy_tile_row( + self, + mSFA_mkl: cute.Tensor, + sSFA: cute.Tensor, + sf_stage, + sfa_c_tile_idx, + sfa_row_off: cutlass.Int32, + sfa_pred_val: cutlass.Boolean, + tidx, + input_C: cutlass.Int32, + ): + sfa_atom_copy = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(), + mSFA_mkl.element_type, + num_bits_per_copy=32, + ) + # The whole 4-warp DMA group cooperates on the SFA copy: each of the 128 + # lanes owns exactly one M token of the 128 x (K/sf_vec_size) tile, so a + # lane loads sf_k_per_ktile scale factors (vec16 -> 8, vec32 -> 4) as + # sf_k_per_ktile//4 chunked 4-byte cp.async copies. The caller hands each + # lane its gmem row offset and predicate, maintained as running deltas + # over the filter walk. + # (warps 8..11 -> local_tid 0..127 via tidx % 128). + local_tid = tidx % 128 + sfa_predicate_tensor = cute.make_rmem_tensor( + cute.make_layout((1,)), + cutlass.Boolean, + ) + sfa_predicate_tensor[0] = sfa_pred_val + # A CTA loads one full SF K tile (self.sf_tile_k channels of scale + # factors) per lane, which the SMEM BlockScaledBasicChunk atom holds. When + # the SF K tile covers several A/B K tiles (ab_k_tiles_per_sf_k_tile > 1), + # the paired A/B tiles share the same filter position, so their SF are + # contiguous in the unswizzled gmem and one SF K tile fills them all; the + # consumer offsets into the live A/B-tile half via sf_base. + sf_k_per_ktile = self.sf_tile_k // self.sf_vec_size + # SFA SMEM uses BlockScaledBasicChunks of blk_sf=4 consecutive scale + # factors (one per channel group). Within a chunk the SF for consecutive + # channel groups are contiguous (identity layout) in both the unswizzled + # global source and the SMEM atom, so a chunk is a single 4-byte + # cp.async; consecutive chunks are separated by the K basic-block stride + # in SMEM, so we issue one copy per chunk. The chunk count adapts to + # sf_vec_size (vec16 -> 8 SF -> 2 chunks; vec32 -> 4 SF -> 1 chunk). + sf_blk = 4 + sf_chunk_cnt = sf_k_per_ktile // sf_blk + sf_k_base = (sfa_c_tile_idx // self.ab_k_tiles_per_sf_k_tile) * sf_k_per_ktile + + # SMEM M coordinate is the 2-level ((32, 4)) mode; the global token of a + # lane is (inner + 32 * outer) == local_tid, matching the layout that the + # SM120 math path reads back. + sf_m_inner = local_tid % 32 + sf_m_outer = local_tid // 32 + + # A pixel owns only ceil(C / sf_vec_size) real scale factors, so a K tile + # wider than C has trailing chunks with no backing storage. Reading one + # would walk into the next pixel's scale factors, and off the end of the + # buffer on the last pixel. Those chunks pair with the A channels past C, + # which the im2col TMA zero-fills, so their scale factors never reach the + # result -- but they must still name an in-bounds address, and must land a + # finite value in SMEM, because the MMA applies the scale before it can + # know the operand is zero. Re-reading the pixel's first chunk gives both. + sf_real_blocks = cute.ceil_div(input_C, self.sf_vec_size) + + for c_chunk in cutlass.range_constexpr(sf_chunk_cnt): + local_sf_k = sf_k_base + c_chunk * sf_blk + gmem_sf_k = local_sf_k if local_sf_k < sf_real_blocks else 0 + tAgSFA_slice = cute.make_tensor( + mSFA_mkl.iterator + (sfa_row_off + gmem_sf_k), + layout=cute.make_layout((sf_blk,)), + ) + # Slice the whole chunk: span the K basic-block (mma_nsf) and the + # blk_sf//mma_nsf middle mode for this chunk c_chunk. Their combined + # extent is always blk_sf and is contiguous from the chunk base, so a + # flat (blk_sf,) layout over the slice iterator is the identity copy. + tAsSFA_slot = sSFA[ + ( + ( + (sf_m_inner, sf_m_outer), + 0, + ), + ((0, None), None, c_chunk), + sf_stage, + ) + ] + tAsSFA_slice = cute.make_tensor( + tAsSFA_slot.iterator, + layout=cute.make_layout((sf_blk,)), + ) + + cute.copy_atom_call( + sfa_atom_copy, + tAgSFA_slice, + tAsSFA_slice, + pred=sfa_predicate_tensor, + ) + + # GPU device kernel + @cute.kernel + def kernel( + self, + tma_atom_a: cute.CopyAtom, + mA_mkl: cute.Tensor, + tma_atom_b: cute.CopyAtom, + mB_nkl: cute.Tensor, + mSFA_mkl: cute.Tensor, + tma_atom_sfb: cute.CopyAtom, + mSFB_nkl: cute.Tensor, + tma_atom_d: cute.CopyAtom, + mD_mnl: cute.Tensor, + tma_atom_residual: Optional[cute.CopyAtom], + mResidual_mnl: Optional[cute.Tensor], + mBias_mnl: Optional[cute.Tensor], + alpha: cutlass.Float32, + beta: cutlass.Constexpr, + tiled_mma: cute.TiledMma, + cta_layout_mnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + sfa_smem_layout_staged: cute.Layout, + sfb_smem_layout_staged: cute.Layout, + sfb_mma_layout_staged: cute.Layout, + epi_smem_layout_staged: cute.ComposedLayout, + res_smem_layout_staged: cute.ComposedLayout, + tile_sched_params: ClcDynamicPersistentTileSchedulerParams, + input_N: cutlass.Int32, + input_D: cutlass.Int32, + input_H: cutlass.Int32, + input_W: cutlass.Int32, + input_C: cutlass.Int32, + output_Z: cutlass.Int32, + output_P: cutlass.Int32, + output_Q: cutlass.Int32, + rt_stride_d: cutlass.Int32, + rt_stride_h: cutlass.Int32, + rt_stride_w: cutlass.Int32, + rt_lower_pad_d: cutlass.Int32, + rt_lower_pad_h: cutlass.Int32, + rt_lower_pad_w: cutlass.Int32, + rt_dil_d: cutlass.Int32, + rt_dil_h: cutlass.Int32, + rt_dil_w: cutlass.Int32, + rt_filter_t: cutlass.Int32, + rt_filter_r: cutlass.Int32, + rt_filter_s: cutlass.Int32, + epilogue_op: cutlass.Constexpr, + ): + """ + GPU device kernel performing the batched GEMM computation. + + :param tma_atom_a: TMA copy atom for A tensor + :type tma_atom_a: cute.CopyAtom + :param mA_mkl: Input tensor A + :type mA_mkl: cute.Tensor + :param tma_atom_b: TMA copy atom for B tensor + :type tma_atom_b: cute.CopyAtom + :param mB_nkl: Input tensor B + :type mB_nkl: cute.Tensor + :param tma_atom_d: TMA copy atom for C tensor + :type tma_atom_d: cute.CopyAtom + :param mD_mnl: Output tensor D + :type mD_mnl: cute.Tensor + :param tiled_mma: Tiled MMA object + :type tiled_mma: cute.TiledMma + :param cta_layout_mnk: CTA layout + :type cta_layout_mnk: cute.Layout + :param a_smem_layout_staged: Shared memory layout for A + :type a_smem_layout_staged: cute.ComposedLayout + :param b_smem_layout_staged: Shared memory layout for B + :type b_smem_layout_staged: cute.ComposedLayout + :param epi_smem_layout_staged: Shared memory layout for epilogue + :type epi_smem_layout_staged: cute.ComposedLayout + """ + + # /////////////////////////////////////////////////////////////////////////////// + # Get cta/warp/thread idx + # /////////////////////////////////////////////////////////////////////////////// + tidx, _, _ = cute.arch.thread_idx() + warp_idx = cute.arch.warp_idx() + warp_idx = cute.arch.make_warp_uniform(warp_idx) + bidx, bidy, bidz = cute.arch.block_idx() + + # ///////////////////////////////////////////////////////////////////////////// + # Prefetch Tma desc + # ///////////////////////////////////////////////////////////////////////////// + if warp_idx == 0: + cpasync.prefetch_descriptor(tma_atom_a) + cpasync.prefetch_descriptor(tma_atom_b) + cpasync.prefetch_descriptor(tma_atom_sfb) + cpasync.prefetch_descriptor(tma_atom_d) + + cta_rank_in_cluster = cute.arch.make_warp_uniform( + cute.arch.block_idx_in_cluster() + ) + cluster_coord_mnk = cta_layout_mnk.get_flat_coord(cta_rank_in_cluster) + + if cutlass.const_expr(cute.rank(a_smem_layout_staged) == 4): + a_smem_layout = cute.slice_(a_smem_layout_staged, (None, None, None, 0)) + else: + a_smem_layout = cute.slice_(a_smem_layout_staged, (None, None, 0)) + if cutlass.const_expr(cute.rank(b_smem_layout_staged) == 4): + b_smem_layout = cute.slice_(b_smem_layout_staged, (None, None, None, 0)) + else: + b_smem_layout = cute.slice_(b_smem_layout_staged, (None, None, 0)) + sfb_smem_layout = cute.slice_(sfb_smem_layout_staged, (None, None, 0)) + tma_copy_bytes = 0 + tma_copy_bytes += cute.size_in_bytes(self.a_dtype, a_smem_layout) + tma_copy_bytes += cute.size_in_bytes(self.b_dtype, b_smem_layout) + tma_copy_bytes += cute.size_in_bytes(self.sf_dtype, sfb_smem_layout) + + # ///////////////////////////////////////////////////////////////////////////// + # Alloc and init AB full/empty + ACC full mbar (pipeline) + # ///////////////////////////////////////////////////////////////////////////// + smem = SmemAllocator() + storage = smem.allocate(self.shared_storage) + + # mbar arrays + mainloop_pipeline_array_ptr = storage.mainloop_pipeline_array_ptr.data_ptr() + math_wg_order_barrier_array_ptr = ( + storage.math_wg_order_barrier_array_ptr.data_ptr() + ) + clc_mbar_ptr = storage.clc_mbar_ptr.data_ptr() + clc_response_ptr = storage.clc_response.data_ptr() + + # Threads/warps participating in this pipeline. One elected TMA producer + # arrive plus one cp.async mbarrier arrive from each SFA cp.async lane. + # The SFA copy is spread across the full DMA warp group (4 warps = 128 + # lanes), so the full barrier expects 128 cp.async arrives plus the + # single TMA producer arrive. + mainloop_pipeline_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.sfa_cpasync_warp_group_size * self.num_threads_per_warp + 1, + ) + mainloop_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, self.num_mma_warps // 2 + ) + + cta_layout_vmnk = cute.make_layout((1, *cta_layout_mnk.shape)) + mainloop_pipeline = pipeline.PipelineTmaAsync.create( + num_stages=self.ab_stage, + producer_group=mainloop_pipeline_producer_group, + consumer_group=mainloop_pipeline_consumer_group, + tx_count=tma_copy_bytes, + barrier_storage=mainloop_pipeline_array_ptr, + cta_layout_vmnk=cta_layout_vmnk, + ) + warp_group_idx = cute.arch.make_warp_uniform(tidx // 128) + math_wg_order_barrier = self.make_and_init_order_barrier( + math_wg_order_barrier_array_ptr, warp_group_idx + ) + + # Residual G2S load pipeline. The dedicated residual warp issues the TMA + # load (single elected producer) into sResidual; the owning math + # warpgroup consumes it (one full-barrier wait per epilogue). Under + # ping-pong only one warpgroup consumes a given tile, so the consumer + # count is a single warpgroup (num_mma_warps // 2 warps). Gated on + # has_residual so beta == 0 allocates nothing. + if cutlass.const_expr(self.has_residual): + residual_pipeline_array_ptr = storage.residual_pipeline_array_ptr.data_ptr() + residual_pipeline_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread + ) + residual_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.num_mma_warps // 2, + ) + residual_pipeline = pipeline.PipelineTmaAsync.create( + num_stages=self.res_stage, + producer_group=residual_pipeline_producer_group, + consumer_group=residual_pipeline_consumer_group, + tx_count=cute.size_in_bytes( + self.d_dtype, + cute.slice_(res_smem_layout_staged, (None, None, 0)), + ), + barrier_storage=residual_pipeline_array_ptr, + cta_layout_vmnk=cta_layout_vmnk, + ) + else: + residual_pipeline = None + + # Bias staging pipeline: the TMA load warp cp.async's each tile's bias + # row into the smem ring and the tile's owning math warpgroup reads it + # back. The producer commit is a per-lane cp.async arrive from that + # warp's 32 threads; the release is a per-thread arrive from the owning + # warpgroup's 128 threads, so each stage's participant set is fixed no + # matter which warpgroup owns the tile. + if cutlass.const_expr(mBias_mnl is not None): + bias_pipeline = pipeline.PipelineCpAsync.create( + num_stages=self.bias_stage, + producer_group=pipeline.CooperativeGroup( + pipeline.Agent.Thread, self.num_threads_per_warp + ), + consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 128), + barrier_storage=storage.bias_pipeline_array_ptr.data_ptr(), + defer_sync=True, + ) + else: + bias_pipeline = None + + # CLC (Cluster Launch Control) pipeline. The scheduler warp is the sole + # producer; the consumers are every math lane, the whole SFA cp.async DMA + # warp group, and the scheduler warp itself (self-consume for exit + # detection). Both math warpgroups peek every tile the scheduler emits and + # a parity counter decides which one owns it, so the consumer count spans + # all math warps, not half. + clc_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + # The residual load warp also walks the tile stream as a CLC consumer when + # beta != 0, so its 32 lanes join the consumer count. This term MUST gate + # on the same const (self.has_residual) as the residual warp's CLC-consume + # block, or the CLC producer/consumer counts diverge and the fetch pipeline + # hangs. + num_residual_clc_warps = 1 if cutlass.const_expr(self.has_residual) else 0 + clc_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + ( + self.num_mma_warps + + self.sfa_cpasync_warp_group_size + + 1 # dedicated TMA load warp + + self.num_sched_warps + + num_residual_clc_warps + ) + * self.num_threads_per_warp, + ) + clc_pipeline = pipeline.PipelineClcFetchAsync.create( + num_stages=self.num_clc_stage, + producer_group=clc_producer_group, + consumer_group=clc_consumer_group, + tx_count=16, + barrier_storage=clc_mbar_ptr, + cta_layout_vmnk=cta_layout_vmnk, + defer_sync=True, + ) + + cute.arch.mbarrier_init_fence() + math_wg_order_state = math_wg_order_barrier.state + + # Cluster arrive after barrier init + if cute.size(self.cluster_shape_mnk) > 1: + cute.arch.cluster_arrive_relaxed() + + # /////////////////////////////////////////////////////////////////////////////// + # Generate smem tensor A/B + # /////////////////////////////////////////////////////////////////////////////// + sA = storage.sA.get_tensor( + a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner + ) + sB = storage.sB.get_tensor( + b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner + ) + sD = storage.sD.get_tensor( + epi_smem_layout_staged.outer, swizzle=epi_smem_layout_staged.inner + ) + # Residual staging uses the same per-subtile atom as the D store so its + # S2R read lines up with the accumulator fragment; it is single-buffered. + if cutlass.const_expr(self.has_residual): + sResidual = storage.sResidual.get_tensor( + res_smem_layout_staged.outer, swizzle=res_smem_layout_staged.inner + ) + else: + sResidual = None + sSFA = storage.sSFA.get_tensor(sfa_smem_layout_staged) + sSFB = storage.sSFB.get_tensor(sfb_smem_layout_staged) + # Same SMEM as sSFB, viewed tile_n wide for the mma consumers. + sSFB_mma = storage.sSFB.get_tensor(sfb_mma_layout_staged) + # (cta_tile_n, STAGE) ring of bias rows: the TMA load warp fills the + # stage its producer state points at, the tile's owning math warpgroup + # reads the stage its consumer state points at. + if cutlass.const_expr(mBias_mnl is not None): + cta_n = self.tile_shape_mnk[1] + sBias = storage.sBias.get_tensor( + cute.make_layout((cta_n, self.bias_stage), stride=(1, cta_n)) + ) + else: + sBias = None + + # /////////////////////////////////////////////////////////////////////////////// + # Local_tile partition global tensors + # /////////////////////////////////////////////////////////////////////////////// + # (bM, bK, loopM, loopK, loopL) + gA_mkl = cute.local_tile( + mA_mkl, + (self.tile_shape_mnk[0], (self.tile_shape_mnk[2], 1, 1, 1)), + (None, None, None), + ) + # (bN, bK, loopN, loopK, loopL). K tile is nested (tileK,) so it tiles + # only the channel (first) level of mB_nkl's (C,S,R,T) K group, leaving + # runtime S/R/T in the loop-K rest — the TMA partition then sees a static + # first mode and takes C and S/R/T at runtime. A flat tileK would + # keep the whole runtime K nest static-checked and fail tma_partition. + gB_nkl = cute.local_tile( + mB_nkl, + (self.tile_shape_mnk[1], (self.tile_shape_mnk[2],)), + (None, None, None), + ) + # (tN, tK, loopN, loopK, loopL). Tiled by the SF K tile (128 for MXFP8), + # so one gmem K tile is a full BlockScaledBasicChunk covering the paired + # A/B K tiles. + gSFB_nkl = cute.local_tile( + mSFB_nkl, + (self.sf_tile_shape_mnk[1], self.sf_tile_shape_mnk[2]), + (None, None, None), + ) + # (bM, bN, loopM, loopN, loopL) + gD_mnl = cute.local_tile( + mD_mnl, + cute.slice_(self.tile_shape_mnk, (None, None, 0)), + (None, None, None), + ) + # (bM, bN, loopM, loopN, loopL) - residual shares the output's MNL tiling + # (same (N,Z,P,Q,K) shape), so it tiles exactly like gD and the G2S load + # lands in the same epi-subtile grid the epilogue consumes. + if cutlass.const_expr(self.has_residual): + gResidual_mnl = cute.local_tile( + mResidual_mnl, + cute.slice_(self.tile_shape_mnk, (None, None, 0)), + (None, None, None), + ) + else: + gResidual_mnl = None + # (bM, bN, loopM, loopN, loopL) - bias shares the output's MNL tiling; the + # M axis carries a stride-0 broadcast so every spatial output row reads the + # same per-output-channel (N) bias. + if cutlass.const_expr(mBias_mnl is not None): + gBias_mnl = cute.local_tile( + mBias_mnl, + cute.slice_(self.tile_shape_mnk, (None, None, 0)), + (None, None, None), + ) + else: + gBias_mnl = None + + # ////////////////////////////////////////////////////////////////////////////// + # Partition global tensor for TiledMMA_A/B/C + # ////////////////////////////////////////////////////////////////////////////// + thr_mma = tiled_mma.get_slice(tidx % 128) + + # ////////////////////////////////////////////////////////////////////////////// + # Partition shared tensor for TMA load A/B + # ////////////////////////////////////////////////////////////////////////////// + # TMA load A partition_S/D + a_cta_layout = cute.make_layout(cute.slice_(cta_layout_mnk, (0, None, 0)).shape) + a_cta_crd = cluster_coord_mnk[1] + if cutlass.const_expr(cute.rank(sA) == 4): + sA_for_tma = cute.group_modes(sA, 0, 3) + else: + sA_for_tma = cute.group_modes(sA, 0, 2) + tAsA, tAgA = cpasync.tma_partition( + tma_atom_a, + a_cta_crd, + a_cta_layout, + sA_for_tma, + cute.group_modes(gA_mkl, 0, 2), + ) + tAsA = cute.filter_zeros(tAsA) + tAgA = cute.filter_zeros(tAgA) + + # TMA load B partition_S/D + b_cta_layout = cute.make_layout(cute.slice_(cta_layout_mnk, (None, 0, 0)).shape) + b_cta_crd = cluster_coord_mnk[0] + if cutlass.const_expr(cute.rank(sB) == 4): + sB_for_tma = cute.group_modes(sB, 0, 3) + else: + sB_for_tma = cute.group_modes(sB, 0, 2) + tBsB, tBgB = cpasync.tma_partition( + tma_atom_b, + b_cta_crd, + b_cta_layout, + sB_for_tma, + cute.group_modes(gB_nkl, 0, 2), + ) + tBsB = cute.filter_zeros(tBsB) + tBgB = cute.filter_zeros(tBgB) + + tBsSFB, tBgSFB = cpasync.tma_partition( + tma_atom_sfb, + b_cta_crd, + b_cta_layout, + cute.group_modes(sSFB, 0, 2), + cute.group_modes(gSFB_nkl, 0, 2), + ) + tBsSFB = cute.filter_zeros(tBsSFB) + tBgSFB = cute.filter_zeros(tBgSFB) + + # Make frangments + tCsA = thr_mma.partition_A(sA) + tCsB = thr_mma.partition_B(sB) + + tCrA = tiled_mma.make_fragment_A(tCsA[None, None, None, 0]) + tCrB = tiled_mma.make_fragment_B(tCsB[None, None, None, 0]) + tCrSFA = sm120_utils.partition_fragment_SFA(sSFA[None, None, 0], thr_mma, tidx) + tCrSFB = sm120_utils.partition_fragment_SFB( + sSFB_mma[None, None, 0], thr_mma, tidx + ) + + tDgD = thr_mma.partition_C(gD_mnl) + acc_shape = tDgD.shape[:3] + accumulators = cute.make_rmem_tensor(acc_shape, self.acc_dtype) + # Bias: same partition_C chain as the accumulator, so the bias fragment + # read back from smem lines up element-for-element with the acc fragment. + # The M-axis stride-0 broadcast (built into mBias_mnl) means each thread + # reads only its own N-column bias. The row is staged into this + # warpgroup's private sBias via cp.async, then loaded to + # registers. + if cutlass.const_expr(mBias_mnl is not None): + # Identity coordinates through the mma C partition, so each fragment + # element carries its global (m, n). The n coordinate (minus the tile's + # n_base) gives the linear column into this warpgroup's sBias row, and + # gates the N-overhang (n >= gemm_n) mask. + cBias_mnl = cute.make_identity_tensor(gBias_mnl.shape) + tCcBias = thr_mma.partition_C(cBias_mnl) + # cp.async transfers at least 32 bits, so each lane moves a 32-bit + # vector of bias elements: two at the output's 16-bit width. + bias_elems_per_copy = 32 // mBias_mnl.element_type.width + bias_g2s_atom = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(), + mBias_mnl.element_type, + num_bits_per_copy=32, + ) + else: + tCcBias = None + bias_elems_per_copy = None + bias_g2s_atom = None + + # cluster wait for barrier init + if cute.size(self.cluster_shape_mnk) > 1: + cute.arch.cluster_wait() + else: + cute.arch.sync_threads() + + k_tile_cnt = cute.size(gA_mkl, mode=[3]) + + # Create the tile scheduler (CLC-based dynamic persistent scheduler). + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + + # Create the pipeline states for producer and consumer + mainloop_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.ab_stage + ) + mainloop_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.ab_stage + ) + + clc_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_clc_stage + ) + + # Residual G2S consumer state. Both math warpgroups advance it in phase: + # the owner consumes+releases each subtile it reads, the non-owner skips + # its tile by advancing the same number of subtile stages the residual + # warp produced for it, so the ring counter stays aligned across tiles. + if cutlass.const_expr(self.has_residual): + residual_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.res_stage + ) + + # Bias ring consumer state, advanced the same way: the TMA load warp + # produces one stage per tile, so the non-owner advances past a skipped + # tile's stage to keep the ring counter aligned with its next owned tile. + if cutlass.const_expr(self.has_bias): + bias_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.bias_stage + ) + + # SFA cp.async warp group + if ( + warp_idx >= self.sfa_cpasync_warp_group_start + and warp_idx < self.tma_load_warp_id + ): + cute.arch.setmaxregister_decrease(self.load_register_requirement) + # The 4-warp SFA group cooperatively issues the SFA cp.async copy + # (one M token per lane). It waits on the empty barrier, loads its + # scale factors, and arrives on the full barrier via the cp.async + # mbarrier path. The dedicated TMA warp (separate warpgroup below) + # issues A/B/SFB and contributes the single TMA arrive to the same + # full barrier. The scalar coordinate iteration is deterministic, so + # running it redundantly on every SFA warp keeps them in sync. + # + # NVFP4 walks the k tiles on hoisted per-tile state: a lane's M + # token decomposition is loop-invariant, so its divides run once per + # output tile, and the row offset and predicate are rebuilt only + # when a filter counter carries. Only the finished offset and + # predicate are read every iteration, which keeps the loop's + # per-lane live set inside the SFA warp register budget. MXFP8 + # keeps the self-contained per-k-tile form. + if cutlass.const_expr(self.sf_vec_size == 16): + sfa_local_tid = tidx % 128 + sfa_ZPQ = output_Z * output_P * output_Q + sfa_PQ = output_P * output_Q + sfa_DHW = input_D * input_H * input_W + sfa_HW = input_H * input_W + sfa_sf_k_extent = mSFA_mkl.shape[1] + while work_tile.is_valid_tile: + tile_coord_mnl = work_tile.tile_idx + + mainloop_producer_state.reset_count() + sfa_c_tile_idx = cutlass.Int32(0) + sfa_s_idx = cutlass.Int32(0) + sfa_r_idx = cutlass.Int32(0) + sfa_t_idx = cutlass.Int32(0) + c_tiles_per_fpos = cute.ceil_div(input_C, self.tile_shape_mnk[2]) + if cutlass.const_expr(self.sf_vec_size == 16): + # A lane's M token and its (n, z, p, q) decomposition hold + # for the whole output tile. + sfa_m_global = ( + tile_coord_mnl[0] * self.tile_shape_mnk[0] + sfa_local_tid + ) + sfa_n_idx = sfa_m_global // sfa_ZPQ + sfa_zpq_rem = sfa_m_global % sfa_ZPQ + sfa_z_idx = sfa_zpq_rem // sfa_PQ + sfa_pq_rem = sfa_zpq_rem % sfa_PQ + sfa_p_idx = sfa_pq_rem // output_Q + sfa_q_idx = sfa_pq_rem % output_Q + sfa_m_valid = sfa_m_global < (input_N * sfa_ZPQ) + sfa_n_clamped = sfa_n_idx if sfa_m_valid else 0 + sfa_d_in = sfa_z_idx * rt_stride_d - rt_lower_pad_d + sfa_h_in = sfa_p_idx * rt_stride_h - rt_lower_pad_h + sfa_w_in = sfa_q_idx * rt_stride_w - rt_lower_pad_w + sfa_m_base = ( + sfa_n_clamped * sfa_DHW + + sfa_d_in * sfa_HW + + sfa_h_in * input_W + + sfa_w_in + ) * sfa_sf_k_extent + # The h and w bases fit 16 bits each (spatial extents are + # far below 32k and the pads below 32), so they ride one + # register; the carry rebuild recovers them with arithmetic + # shifts, keeping the loop's carried live set minimal. + sfa_d_base = sfa_d_in + 0 + sfa_hw_base = (sfa_h_in << 16) | (sfa_w_in & 0xFFFF) + # The offset and predicate change only when a filter counter + # does, so they are computed here and after each carry; a + # predicated-off row reads address 0, always in bounds, and + # its predicate suppresses the actual copy. + sfa_pred_val = cutlass.Boolean(0) + if sfa_m_valid: + if sfa_d_in >= 0 and sfa_d_in < input_D: + if sfa_h_in >= 0 and sfa_h_in < input_H: + if sfa_w_in >= 0 and sfa_w_in < input_W: + sfa_pred_val = cutlass.Boolean(1) + sfa_row_off = sfa_m_base if sfa_pred_val else 0 + + # Counting up would keep the k-tile bound live across the + # barrier wait below and spill it to the stack. + sfa_iters_left = k_tile_cnt + 0 + while sfa_iters_left > 0: + # Wait for the A/B/SFA buffers to be empty before + # writing. + mainloop_pipeline.sync_object_empty.wait( + mainloop_producer_state.index, + mainloop_producer_state.phase, + ) + # All SFA lanes cooperatively load the SFA tile and + # each arrive on the full barrier when their cp.async + # completes. + self.sfa_cpasync_copy_tile_row( + mSFA_mkl, + sSFA, + mainloop_producer_state.index, + sfa_c_tile_idx, + sfa_row_off, + sfa_pred_val, + tidx, + input_C, + ) + cute.arch.cp_async_commit_group() + mainloop_pipeline.sync_object_full.arrive_cp_async_mbarrier( + mainloop_producer_state.index + ) + mainloop_producer_state.advance() + sfa_c_tile_idx = sfa_c_tile_idx + 1 + if sfa_c_tile_idx >= c_tiles_per_fpos: + sfa_c_tile_idx = 0 + sfa_s_idx = sfa_s_idx + 1 + if sfa_s_idx >= rt_filter_s: + sfa_s_idx = 0 + sfa_r_idx = sfa_r_idx + 1 + if sfa_r_idx >= rt_filter_r: + sfa_r_idx = 0 + sfa_t_idx = sfa_t_idx + 1 + if sfa_t_idx >= rt_filter_t: + sfa_t_idx = 0 + # Any carry level may have wrapped, so rebuild the + # coords, offset and predicate from the counters; + # this runs once per filter position, not per tile. + sfa_d_in = sfa_d_base + sfa_t_idx * rt_dil_d + sfa_h_in = (sfa_hw_base >> 16) + sfa_r_idx * rt_dil_h + sfa_w_in = ( + (sfa_hw_base << 16) >> 16 + ) + sfa_s_idx * rt_dil_w + sfa_row_delta = ( + sfa_t_idx * rt_dil_d * sfa_HW + + sfa_r_idx * rt_dil_h * input_W + + sfa_s_idx * rt_dil_w + ) * sfa_sf_k_extent + sfa_pred_val = cutlass.Boolean(0) + if sfa_m_valid: + if sfa_d_in >= 0 and sfa_d_in < input_D: + if sfa_h_in >= 0 and sfa_h_in < input_H: + if sfa_w_in >= 0 and sfa_w_in < input_W: + sfa_pred_val = cutlass.Boolean(1) + sfa_row_off = ( + (sfa_m_base + sfa_row_delta) if sfa_pred_val else 0 + ) + sfa_iters_left -= 1 + else: + for k_tile in range(0, k_tile_cnt, 1, unroll=1): + # Wait for the A/B/SFA buffers to be empty before + # writing. + mainloop_pipeline.sync_object_empty.wait( + mainloop_producer_state.index, + mainloop_producer_state.phase, + ) + # All SFA lanes cooperatively load the SFA tile and + # each arrive on the full barrier when their cp.async + # completes. + self.sfa_cpasync_copy_tile( + mSFA_mkl, + sSFA, + tile_coord_mnl, + mainloop_producer_state.index, + sfa_c_tile_idx, + sfa_s_idx, + sfa_r_idx, + sfa_t_idx, + tidx, + input_N, + input_C, + input_D, + input_H, + input_W, + output_Z, + output_P, + output_Q, + rt_stride_d, + rt_stride_h, + rt_stride_w, + rt_lower_pad_d, + rt_lower_pad_h, + rt_lower_pad_w, + rt_dil_d, + rt_dil_h, + rt_dil_w, + ) + cute.arch.cp_async_commit_group() + mainloop_pipeline.sync_object_full.arrive_cp_async_mbarrier( + mainloop_producer_state.index + ) + mainloop_producer_state.advance() + sfa_c_tile_idx = sfa_c_tile_idx + 1 + if sfa_c_tile_idx >= c_tiles_per_fpos: + sfa_c_tile_idx = 0 + sfa_s_idx = sfa_s_idx + 1 + if sfa_s_idx >= rt_filter_s: + sfa_s_idx = 0 + sfa_r_idx = sfa_r_idx + 1 + if sfa_r_idx >= rt_filter_r: + sfa_r_idx = 0 + sfa_t_idx = sfa_t_idx + 1 + if sfa_t_idx >= rt_filter_t: + sfa_t_idx = 0 + + # Pull the next tile from the CLC response slot. The whole SFA + # cp.async warp group participates as a CLC consumer. + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + # end of while loop + + # Wait A/B/SFA buffers empty before the SFA warps exit (all lanes + # only wait here, so the whole group can participate safely). + mainloop_pipeline.producer_tail(mainloop_producer_state) + # Dedicated TMA load warp, carried in the CLC warpgroup + elif warp_idx == self.tma_load_warp_id: + cute.arch.setmaxregister_decrease(self.sched_register_requirement) + # Sole issuer of the A/B/SFB TMA loads. Sets the transaction barrier + # for each stage and contributes the single TMA arrive; the SFA + # warpgroup contributes the cp.async arrives to the same full + # barrier. Walks the same tile stream as a CLC consumer. + if cutlass.const_expr(mBias_mnl is not None): + bias_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.bias_stage + ) + while work_tile.is_valid_tile: + tile_coord_mnl = work_tile.tile_idx + tAgA_mkl = tAgA[(None, tile_coord_mnl[0], None, tile_coord_mnl[2])] + tBgB_nkl = tBgB[(None, tile_coord_mnl[1], None, tile_coord_mnl[2])] + # gSFB is tiled by the 128-wide SF chunk, so when tile_n is + # narrower than a chunk several B N tiles map onto the same SF + # tile and the SF N coordinate scales down accordingly. + sf_n_coord = tile_coord_mnl[1] // self.n_tiles_per_sf_n_tile + tBgSFB_nkl = tBgSFB[(None, sf_n_coord, None, tile_coord_mnl[2])] + + mainloop_producer_state.reset_count() + k_shape = cute.shape(tAgA_mkl, mode=1) + coord_iter = cute.repeat_like(0, k_shape) + sfa_c_tile_idx = cutlass.Int32(0) + sfa_s_idx = cutlass.Int32(0) + sfa_r_idx = cutlass.Int32(0) + sfa_t_idx = cutlass.Int32(0) + c_tiles_per_fpos = cute.ceil_div(input_C, self.tile_shape_mnk[2]) + # SF K tiles one filter position spans, the count the host padded + # SFB's per-position span to. + sfb_c_tiles = cute.ceil_div(input_C, self.sf_tile_k) + + for k_tile in range(0, k_tile_cnt, 1, unroll=1): + # Wait for the A/B/SFA buffers to be empty before writing. + mainloop_pipeline.sync_object_empty.wait( + mainloop_producer_state.index, mainloop_producer_state.phase + ) + # Set the transaction barrier for the A/B/SFB buffers, then + # issue the TMA loads. + mainloop_pipeline.sync_object_full.arrive( + mainloop_producer_state.index, + mainloop_pipeline.producer_mask, + ) + + tAgA_k = tAgA_mkl[(None, coord_iter)] + tAsA_pipe = tAsA[(None, mainloop_producer_state.index)] + + tBgB_k = tBgB_nkl[ + ( + None, + self._make_b_k_coord( + sfa_c_tile_idx, + sfa_s_idx, + sfa_r_idx, + sfa_t_idx, + ), + ) + ] + tBsB_pipe = tBsB[(None, mainloop_producer_state.index)] + + # SFB's K mode is flat: it runs one filter position's channels + # before moving to the next, and the positions in the S->R->T + # order A and B traverse. The span padding puts each position on + # a tile boundary, so a position holds sfb_c_tiles whole tiles + # and a tile's flat index is that count times the position plus + # the channel tile inside it. The channel tile is on the SF + # cadence, which is coarser whenever one SF chunk serves several + # A/B K tiles; each A/B stage still TMAs the whole chunk and the + # consumer reads its live part. + sfb_fpos = ( + sfa_s_idx + + sfa_r_idx * rt_filter_s + + sfa_t_idx * (rt_filter_r * rt_filter_s) + ) + tBgSFB_k = tBgSFB_nkl[ + ( + None, + sfb_fpos * sfb_c_tiles + + sfa_c_tile_idx // self.ab_k_tiles_per_sf_k_tile, + ) + ] + tBsSFB_pipe = tBsSFB[(None, mainloop_producer_state.index)] + + cute.copy( + tma_atom_a, + tAgA_k, + tAsA_pipe, + tma_bar_ptr=mainloop_pipeline.producer_get_barrier( + mainloop_producer_state + ), + ) + cute.copy( + tma_atom_b, + tBgB_k, + tBsB_pipe, + tma_bar_ptr=mainloop_pipeline.producer_get_barrier( + mainloop_producer_state + ), + ) + cute.copy( + tma_atom_sfb, + tBgSFB_k, + tBsSFB_pipe, + tma_bar_ptr=mainloop_pipeline.producer_get_barrier( + mainloop_producer_state + ), + ) + # Mainloop pipeline's producer commit is a NOP for TMA. + mainloop_pipeline.producer_commit(mainloop_producer_state) + mainloop_producer_state.advance() + sfa_c_tile_idx = sfa_c_tile_idx + 1 + if sfa_c_tile_idx >= c_tiles_per_fpos: + sfa_c_tile_idx = 0 + sfa_s_idx = sfa_s_idx + 1 + if sfa_s_idx >= rt_filter_s: + sfa_s_idx = 0 + sfa_r_idx = sfa_r_idx + 1 + if sfa_r_idx >= rt_filter_r: + sfa_r_idx = 0 + sfa_t_idx = sfa_t_idx + 1 + if sfa_t_idx >= rt_filter_t: + sfa_t_idx = 0 + coord_iter = cute.increment_coord(coord_iter, k_shape) + + # Stage this tile's bias row for its owning math warpgroup: wait + # for the ring stage to drain, cp.async the row in, and commit so + # each lane's arrive lands once its copies complete. + if cutlass.const_expr(mBias_mnl is not None): + bias_pipeline.producer_acquire(bias_producer_state) + bias_lane_idx = tidx % self.num_threads_per_warp + n_base = tile_coord_mnl[1] * cta_n + # Each lane cp.async's contiguous 32-bit vectors + # (bias_elems_per_copy elements each) of this tile's bias row. + # Column-major stride so lane t owns [t*elems, (t+1)*elems); + # the output-channel axis has stride 1, so each segment is + # contiguous. n_active lanes cover cta_n. + n_active = cta_n // bias_elems_per_copy + bias_row_layout = cute.make_layout( + (bias_elems_per_copy, n_active), + stride=(1, bias_elems_per_copy), + ) + # cp.async needs 32-bit source and destination alignment; the + # tile base is a multiple of cta_tile_n, which keeps both on a + # 4-byte boundary, so re-annotate the pointers to satisfy the + # 32-bit atom. + gBias_row = cute.make_tensor( + (mBias_mnl.iterator + n_base).align(min_align=4), + bias_row_layout, + ) + sBias_stage = sBias[(None, bias_producer_state.index)] + sBias_tiled = cute.make_tensor( + sBias_stage.iterator.align(min_align=4), bias_row_layout + ) + bias_pred = cute.make_rmem_tensor( + cute.make_layout((1,)), cutlass.Boolean + ) + # The warp has 32 lanes and the row needs n_active of them; + # each extra pass hands the lanes the next 32 vectors. + for bias_pass in cutlass.range_constexpr((n_active + 31) // 32): + bias_lane = bias_pass * 32 + bias_lane_idx + if bias_lane < n_active: + # A CTA N-tile rounds up to cta_tile_n, but the + # output-channel count need not divide it, so tail + # lanes address bias columns past the end with no + # backing storage. Guard each lane's vector on its + # base channel: in-bounds lanes copy from gmem, + # out-of-bounds lanes zero-fill (cp.async writes 0 on + # a false predicate). The zero tail is only read back + # for overhang output the TMA store clamps away. + bias_pred[0] = cutlass.Boolean( + n_base + bias_lane * bias_elems_per_copy + < mBias_mnl.shape[1] + ) + cute.copy_atom_call( + bias_g2s_atom, + gBias_row[(None, bias_lane)], + sBias_tiled[(None, bias_lane)], + pred=bias_pred, + ) + cute.arch.cp_async_commit_group() + bias_pipeline.producer_commit(bias_producer_state) + bias_producer_state.advance() + + # Walk the tile stream in lockstep with the SFA group. + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + # end of while loop + + mainloop_pipeline.producer_tail(mainloop_producer_state) + # MMA warp group + elif warp_idx < self.num_mma_warps: + cute.arch.setmaxregister_increase(self.mma_register_requirement) + num_k_blocks = cute.size(tCrA, mode=[2]) + # A/B K tiles one filter position spans. sf_base below indexes the SF + # chunk by the tile's index inside its own filter position, which is what + # stays aligned with the chunk when a position holds a fractional number + # of chunks. + c_tiles_per_fpos = cute.ceil_div(input_C, self.tile_shape_mnk[2]) + + # /////////////////////////////////////////////////////////////////////////////// + # Copy Atom A/B retiling for TMA load A/B + # /////////////////////////////////////////////////////////////////////////////// + atom_copy_ldmatrix_A = make_ldmatrix_atom( + self.a_dtype, + transpose=self.a_layout.is_m_major_a(), + num_matrices=4, + mixed_mode=self.mixed_mode, + ) + atom_copy_ldmatrix_B = make_ldmatrix_atom( + self.b_dtype, + transpose=self.b_layout.is_n_major_b(), + num_matrices=4, + mixed_mode=self.mixed_mode, + ) + smem_tiled_copy_A = cute.make_tiled_copy_A(atom_copy_ldmatrix_A, tiled_mma) + smem_tiled_copy_B = cute.make_tiled_copy_B(atom_copy_ldmatrix_B, tiled_mma) + + atom_copy_ldmatrix_SF = cute.make_copy_atom( + cute.nvgpu.CopyUniversalOp(), + self.sf_dtype, + ) + smem_tiled_copy_SFA = cute.make_tiled_copy( + atom_copy_ldmatrix_SF, + sm120_utils.get_layoutSFA_TV(tiled_mma), + ( + cute.size(tiled_mma.permutation_mnk[0]), + cute.size(tiled_mma.permutation_mnk[2]), + ), + ) + smem_tiled_copy_SFB = cute.make_tiled_copy( + atom_copy_ldmatrix_SF, + sm120_utils.get_layoutSFB_TV(tiled_mma), + ( + cute.size(tiled_mma.permutation_mnk[1]), + cute.size(tiled_mma.permutation_mnk[2]), + ), + ) + + thr_copy_ldmatrix_A = smem_tiled_copy_A.get_slice(tidx % 128) + thr_copy_ldmatrix_B = smem_tiled_copy_B.get_slice(tidx % 128) + tCsA_copy_view = thr_copy_ldmatrix_A.partition_S(sA) + tCrA_copy_view = thr_copy_ldmatrix_A.retile(tCrA) + tCsB_copy_view = thr_copy_ldmatrix_B.partition_S(sB) + tCrB_copy_view = thr_copy_ldmatrix_B.retile(tCrB) + + thr_copy_ldmatrix_SFA = smem_tiled_copy_SFA.get_slice(tidx % 128) + thr_copy_ldmatrix_SFB = smem_tiled_copy_SFB.get_slice(tidx % 128) + tCsSFA_copy_view = thr_copy_ldmatrix_SFA.partition_S(sSFA) + tCrSFA_copy_view = thr_copy_ldmatrix_SFA.retile(tCrSFA) + tCsSFB_copy_view = thr_copy_ldmatrix_SFB.partition_S(sSFB_mma) + tCrSFB_copy_view = thr_copy_ldmatrix_SFB.retile(tCrSFB) + + # Parity counter for the ping-pong warpgroups. The CLC scheduler + # hands the same tile stream to both warpgroups; wg0 owns the tiles + # at parity 0 (tile 0, 2, 4, ...) and wg1 owns parity 1. The + # warpgroup whose parity does NOT match skips the tile and advances + # its mainloop_consumer_state by k_tile_cnt so DMA's linear + # production stays in phase with this warpgroup's next owned tile. + tile_parity = cutlass.Int32(0) + + while work_tile.is_valid_tile: + if tile_parity == warp_group_idx: + tile_coord_mnl = work_tile.tile_idx + gD_mnl_slice = gD_mnl[(None, None, *tile_coord_mnl)] + # This tile's slice of the staged SF chunk. Consecutive N + # tiles share one chunk when tile_n is narrower than it, and + # a slice sits (tile_n/32) SF blocks past the previous one. + if cutlass.const_expr(self.n_tiles_per_sf_n_tile > 1): + sfb_n_slice = tile_coord_mnl[1] % self.n_tiles_per_sf_n_tile + tCsSFB_tile = cute.make_tensor( + tCsSFB_copy_view.iterator + + sfb_n_slice * (self.tile_shape_mnk[1] // 32) * 4, + tCsSFB_copy_view.layout, + ) + else: + tCsSFB_tile = tCsSFB_copy_view + # Clear the accumulator + accumulators.fill(0.0) + + # ///////////////////////////////////////////////////////////////////////////// + # Pipelined MAINLOOP + # ///////////////////////////////////////////////////////////////////////////// + + mainloop_consumer_state.reset_count() + math_wg_order_barrier.wait(math_wg_order_state) + peek_ab_full_status = cutlass.Boolean(1) + if mainloop_consumer_state.count < k_tile_cnt: + peek_ab_full_status = mainloop_pipeline.consumer_try_wait( + mainloop_consumer_state + ) + + # Wait for TMA copies to complete + mainloop_pipeline.consumer_wait( + mainloop_consumer_state, peek_ab_full_status + ) + # tCsA_p: (MMA, (4, MMA_M / 4), MMA_K), tCsA_p: (MMA, (4, MMA_N / 4), MMA_K) + tCsA_p = tCsA_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + tCsB_p = tCsB_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + tCsSFA_p = tCsSFA_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + tCsSFB_p = tCsSFB_tile[ + None, None, None, mainloop_consumer_state.index + ] + cute.copy( + smem_tiled_copy_A, + tCsA_p[None, None, 0], + tCrA_copy_view[None, None, 0], + ) + cute.copy( + smem_tiled_copy_B, + tCsB_p[None, None, 0], + tCrB_copy_view[None, None, 0], + ) + + tCsSFA_p_filtered = cute.filter_zeros(tCsSFA_p) + tCsSFB_p_filtered = cute.filter_zeros(tCsSFB_p) + tCrSFA_copy_view_filtered = cute.filter_zeros(tCrSFA_copy_view) + tCrSFB_copy_view_filtered = cute.filter_zeros(tCrSFB_copy_view) + + # One SF K tile holds ab_k_tiles_per_sf_k_tile A/B tiles worth + # of scale factors. sf_base picks the current A/B tile's half + # of the chunk (0 for NVFP4, where the ratio is 1). + sf_base = ( + (mainloop_consumer_state.count % c_tiles_per_fpos) + % self.ab_k_tiles_per_sf_k_tile + ) * num_k_blocks + cute.copy( + smem_tiled_copy_SFA, + tCsSFA_p_filtered[None, None, sf_base], + tCrSFA_copy_view_filtered[None, None, 0], + ) + cute.copy( + smem_tiled_copy_SFB, + tCsSFB_p_filtered[None, None, sf_base], + tCrSFB_copy_view_filtered[None, None, 0], + ) + + for k_tile in range(0, k_tile_cnt - 1, 1, unroll=1): + # unroll the loop + for k_block_idx in cutlass.range_constexpr(num_k_blocks): + k_block_next = ( + 0 + if k_block_idx + 1 == num_k_blocks + else k_block_idx + 1 + ) + + if k_block_idx == num_k_blocks - 1: + mainloop_pipeline.consumer_release( + mainloop_consumer_state + ) + mainloop_consumer_state.advance() + + peek_ab_full_status = cutlass.Boolean(1) + peek_ab_full_status = ( + mainloop_pipeline.consumer_try_wait( + mainloop_consumer_state + ) + ) + + tCsA_p = tCsA_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + tCsB_p = tCsB_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + tCsSFA_p = tCsSFA_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + tCsSFB_p = tCsSFB_tile[ + None, None, None, mainloop_consumer_state.index + ] + mainloop_pipeline.consumer_wait( + mainloop_consumer_state, peek_ab_full_status + ) + + def prefetch_next_k_block(): + cute.copy( + smem_tiled_copy_A, + tCsA_p[None, None, k_block_next], + tCrA_copy_view[None, None, k_block_next], + ) + cute.copy( + smem_tiled_copy_B, + tCsB_p[None, None, k_block_next], + tCrB_copy_view[None, None, k_block_next], + ) + # sf_base picks the live A/B tile's part of the SF + # K chunk; the consumer count it reads has already + # been advanced above when k_block_idx is the last one. + sf_base = ( + (mainloop_consumer_state.count % c_tiles_per_fpos) + % self.ab_k_tiles_per_sf_k_tile + ) * num_k_blocks + cute.copy( + smem_tiled_copy_SFA, + cute.filter_zeros(tCsSFA_p)[ + None, None, sf_base + k_block_next + ], + cute.filter_zeros(tCrSFA_copy_view)[ + None, None, k_block_next + ], + ) + cute.copy( + smem_tiled_copy_SFB, + cute.filter_zeros(tCsSFB_p)[ + None, None, sf_base + k_block_next + ], + cute.filter_zeros(tCrSFB_copy_view)[ + None, None, k_block_next + ], + ) + + # With more than one k_block the prefetch targets a + # different register slot than the MMA is about to read, + # so issuing it first overlaps the load with the MMA. + # With exactly one, k_block_next wraps to the same slot + # and the tile has already been released, so prefetching + # first would overwrite the operands with the next tile's + # before the MMA consumes them; it has to come after. + if cutlass.const_expr(num_k_blocks > 1): + prefetch_next_k_block() + # Mixed FP4 x FP8 register-side bit shift before mma.sync to + # move the FP4 nibble (loaded into the LOW half of each Int8 + # register byte by ldsm.b4x16_p64) into the MIDDLE of the byte + # where mma.sync.kind::mxf8f6f4 reads it. See blockscaled_gemm_dispatch.FP4_SHIFT_BITS. + if cutlass.const_expr( + self.mixed_mode and self.a_dtype.width < 8 + ): + a_view = cute.recast_tensor( + tCrA[None, None, k_block_idx], cutlass.Int8 + ) + for _i in cutlass.range_constexpr(cute.size(a_view)): + a_view[_i] = cutlass.Int8( + a_view[_i] << FP4_SHIFT_BITS + ) + if cutlass.const_expr( + self.mixed_mode and self.b_dtype.width < 8 + ): + b_view = cute.recast_tensor( + tCrB[None, None, k_block_idx], cutlass.Int8 + ) + for _i in cutlass.range_constexpr(cute.size(b_view)): + b_view[_i] = cutlass.Int8( + b_view[_i] << FP4_SHIFT_BITS + ) + cute.gemm( + tiled_mma, + accumulators, + [ + tCrA[None, None, k_block_idx], + tCrSFA[None, None, k_block_idx], + ], + [ + tCrB[None, None, k_block_idx], + tCrSFB[None, None, k_block_idx], + ], + accumulators, + ) + if cutlass.const_expr(num_k_blocks == 1): + prefetch_next_k_block() + # end of for loop + # end of for loop + # Hoist out last k_tile + for k_block_idx in cutlass.range_constexpr(num_k_blocks): + k_block_next = ( + 0 if k_block_idx + 1 == num_k_blocks else k_block_idx + 1 + ) + + if k_block_idx == num_k_blocks - 1: + mainloop_pipeline.consumer_release(mainloop_consumer_state) + mainloop_consumer_state.advance() + + if k_block_next > 0: + cute.copy( + smem_tiled_copy_A, + tCsA_p[None, None, k_block_next], + tCrA_copy_view[None, None, k_block_next], + ) + cute.copy( + smem_tiled_copy_B, + tCsB_p[None, None, k_block_next], + tCrB_copy_view[None, None, k_block_next], + ) + tCsSFA_p_filtered = cute.filter_zeros(tCsSFA_p) + tCsSFB_p_filtered = cute.filter_zeros(tCsSFB_p) + tCrSFA_copy_view_filtered = cute.filter_zeros( + tCrSFA_copy_view + ) + tCrSFB_copy_view_filtered = cute.filter_zeros( + tCrSFB_copy_view + ) + sf_base = ( + (mainloop_consumer_state.count % c_tiles_per_fpos) + % self.ab_k_tiles_per_sf_k_tile + ) * num_k_blocks + cute.copy( + smem_tiled_copy_SFA, + tCsSFA_p_filtered[None, None, sf_base + k_block_next], + tCrSFA_copy_view_filtered[None, None, k_block_next], + ) + cute.copy( + smem_tiled_copy_SFB, + tCsSFB_p_filtered[None, None, sf_base + k_block_next], + tCrSFB_copy_view_filtered[None, None, k_block_next], + ) + # Mixed FP4 x FP8 register-side bit shift before mma.sync (hoisted tail). + if cutlass.const_expr( + self.mixed_mode and self.a_dtype.width < 8 + ): + a_view_h = cute.recast_tensor( + tCrA[None, None, k_block_idx], cutlass.Int8 + ) + for _i in cutlass.range_constexpr(cute.size(a_view_h)): + a_view_h[_i] = cutlass.Int8( + a_view_h[_i] << FP4_SHIFT_BITS + ) + if cutlass.const_expr( + self.mixed_mode and self.b_dtype.width < 8 + ): + b_view_h = cute.recast_tensor( + tCrB[None, None, k_block_idx], cutlass.Int8 + ) + for _i in cutlass.range_constexpr(cute.size(b_view_h)): + b_view_h[_i] = cutlass.Int8( + b_view_h[_i] << FP4_SHIFT_BITS + ) + cute.gemm( + tiled_mma, + accumulators, + [ + tCrA[None, None, k_block_idx], + tCrSFA[None, None, k_block_idx], + ], + [ + tCrB[None, None, k_block_idx], + tCrSFB[None, None, k_block_idx], + ], + accumulators, + ) + + # Signal the other warpgroup it may start its mainloop. The + # matching wait is deferred to just before the epilogue so the + # two warpgroups overlap this warpgroup's epilogue with the + # other's mainloop (CUTLASS ping-pong). + math_wg_order_state = math_wg_order_barrier.arrive( + math_wg_order_state + ) + + # Apply the per-tensor alpha scale in FP32 before bias/residual + # (D = act(alpha * acc + bias)). alpha is a runtime scalar, so + # this multiply lives in every cubin regardless of its value. + accumulators.store(accumulators.load() * alpha) + + # Add per-output-channel bias in FP32 (D = act(alpha*acc + bias)). + # The TMA load warp has cp.async'd this tile's contiguous + # cta_tile_n bias values into the ring stage this warpgroup's + # consumer state points at; every thread loads its own N + # columns back and adds them to the accumulator. One value per + # output channel is fetched once and broadcast to all M rows + # via the stride-0 M axis of mBias_mnl. + if cutlass.const_expr(mBias_mnl is not None): + bias_pipeline.consumer_wait(bias_consumer_state) + sBias_row = sBias[(None, bias_consumer_state.index)] + # Read back into a fragment aligned with the accumulator. The + # identity-coordinate partition gives each element its tile-local + # N column [0, cta_n), which is exactly the linear index into + # the staged row (the producer already applied n_base when + # staging from gmem and zero-filled any N-overhang). Indexing + # the row by that tile-local column, rather than by a gmem + # channel stride, is what makes the smem read correct. + tCcBias_tile = tCcBias[(None, None, None, *tile_coord_mnl)] + tCrBias = cute.make_rmem_tensor( + accumulators.shape, mBias_mnl.element_type + ) + for be in cutlass.range_constexpr(cute.size(tCrBias)): + tCrBias[be] = sBias_row[tCcBias_tile[be][1]] + # Release the stage so the TMA load warp can refill it: + # every thread arrives only after its smem reads are done. + cute.arch.fence_proxy("async.shared", space="cta") + bias_pipeline.consumer_release(bias_consumer_state) + bias_consumer_state.advance() + accumulators.store( + accumulators.load() + tCrBias.load().to(self.acc_dtype) + ) + + # ///////////////////////////////////////////////////////////////////////////// + # EPILOG + # ///////////////////////////////////////////////////////////////////////////// + + # Serialize the epilogue against the other warpgroup: wait for + # our turn before touching the shared sD staging buffer. Pairs + # with the mainloop-tail arrive above so this warpgroup's + # epilogue overlaps the other's mainloop. + math_wg_order_barrier.wait(math_wg_order_state) + ( + tiled_copy_r2s, + thr_copy_r2s, + tRS_sD, + tRS_rAcc, + tRS_rD, + tRS_rD_layout, + ) = self.epilog_smem_copy_and_partition( + tiled_mma, accumulators, tidx, sD + ) + + bSG_sD, bSG_gD = self.epilog_gmem_copy_and_partition( + tma_atom_d, gD_mnl_slice, sD + ) + + # Residual read-back: partition sResidual with the same thread + # slice the D store uses for its reg source (partition_S), so + # tRS_sResidual has the exact per-thread reg-fragment TV layout + # of tRS_rD (which is allocated from this partition's shape). + # A plain element copy then reads each stage into a congruent + # reg fragment, so the FP32 add lines up with the accumulator. + if cutlass.const_expr(self.has_residual): + # (R2S, R2S_M, R2S_N, PIPE_D) reg-side view of the residual + # smem, congruent with tRS_rD. + tRS_sResidual = thr_copy_r2s.partition_S(sResidual) + + # Initialize tma store pipeline + tma_store_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.num_mma_warps * self.num_threads_per_warp, + ) + tma_store_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.epi_stage, + producer_group=tma_store_producer_group, + ) + + epi_rest_m = bSG_gD.shape[1][0] + epi_rest_n = bSG_gD.shape[1][1] + epi_tile_m = self.epi_tile[0] + epi_tile_n = self.epi_tile[1] + mma_tile_m = self.tile_shape_mnk[0] // cute.size(tRS_rAcc, mode=[1]) + mma_tile_n = self.tile_shape_mnk[1] // cute.size(tRS_rAcc, mode=[2]) + + for epi_m in cutlass.range_constexpr(epi_rest_m): + for epi_n in cutlass.range_constexpr(epi_rest_n): + MmaMPerEpiM = epi_tile_m // mma_tile_m + MmaNPerEpiN = epi_tile_n // mma_tile_n + for mma_n_in_epi in cutlass.range_constexpr(MmaNPerEpiN): + for mma_m_in_epi in cutlass.range_constexpr( + MmaMPerEpiM + ): + mma_n = (epi_n * MmaNPerEpiN) + mma_n_in_epi + mma_m = (epi_m * MmaMPerEpiM) + mma_m_in_epi + tRS_rD_slice = tRS_rD[ + (None, mma_m_in_epi, mma_n_in_epi) + ] + tRS_rAcc_slice = tRS_rAcc[(None, mma_m, mma_n)] + for elem_idx in cutlass.range_constexpr( + cute.size(tRS_rD_slice) + ): + tRS_rD_slice[elem_idx] = tRS_rAcc_slice[ + elem_idx + ] + + # Add this subtile's residual in FP32 before the + # activation (D = act(alpha*acc + bias + beta*residual)). + # Wait for the residual warp's G2S load to land, read it + # into a fragment shaped like tRS_rD, add beta*residual, + # then release the stage. beta is a compile-time constant + # folded into the scaled add. + if cutlass.const_expr(self.has_residual): + residual_pipeline.consumer_wait(residual_consumer_state) + tRS_rResidual = cute.make_rmem_tensor( + tRS_rD_layout.shape, self.d_dtype + ) + cute.autovec_copy( + tRS_sResidual[ + ( + None, + None, + None, + residual_consumer_state.index, + ) + ], + tRS_rResidual, + ) + cute.arch.fence_proxy( + "async.shared", + space="cta", + ) + residual_pipeline.consumer_release( + residual_consumer_state + ) + residual_consumer_state.advance() + tRS_rD.store( + tRS_rD.load() + + self.beta + * tRS_rResidual.load().to(self.acc_dtype) + ) + + # Apply the epilogue activation in FP32, then cast to the + # output type. Folded into the cubin as a Constexpr op, so + # each activation compiles to its own kernel. + tRS_rD_out = cute.make_rmem_tensor( + tRS_rD_layout.shape, self.d_dtype + ) + acc_vec = epilogue_op(tRS_rD.load()) + tRS_rD_out.store(acc_vec.to(self.d_dtype)) + + # Register to shared memory + epi_buffer = (epi_m * epi_rest_n + epi_n) % cute.size( + tRS_sD, mode=[3] + ) + self.epilog_sync_barrier.arrive_and_wait() + cute.copy( + tiled_copy_r2s, + tRS_rD_out, + tRS_sD[(None, None, None, epi_buffer)], + ) + cute.arch.fence_proxy( + "async.shared", + space="cta", + ) + # barrier for sync + self.epilog_sync_barrier.arrive_and_wait() + # Get the global memory coordinate for the current epi tile. + gmem_coord = (epi_m, epi_n) + # Copy from shared memory to global memory + if warp_idx % 4 == 0: + cute.copy( + tma_atom_d, + bSG_sD[(None, epi_buffer)], + bSG_gD[(None, gmem_coord)], + ) + tma_store_pipeline.producer_commit() + tma_store_pipeline.producer_acquire() + tma_store_pipeline.producer_tail() + # Signal the other warpgroup it can start its epilogue. + math_wg_order_state = math_wg_order_barrier.arrive( + math_wg_order_state + ) + else: + # This warpgroup does not own the current tile. Advance the + # mainloop consumer state by the k_tile_cnt stages the DMA + # warp group produced for it, keeping DMA's linear production + # in phase with this warpgroup's next owned tile. + mainloop_consumer_state = self.advance( + mainloop_consumer_state, k_tile_cnt + ) + # Likewise skip this tile's residual: the residual warp + # produced one stage per output subtile for it, so advance the + # residual consumer state by that subtile count to keep the + # ring counter aligned with the next owned tile. + if cutlass.const_expr(self.has_residual): + residual_subtiles_per_tile = ( + self.tile_shape_mnk[0] // self.epi_tile[0] + ) * (self.tile_shape_mnk[1] // self.epi_tile[1]) + residual_consumer_state = self.advance( + residual_consumer_state, residual_subtiles_per_tile + ) + # Likewise skip this tile's bias stage: the TMA load warp + # produced one per tile, so advance past it to stay aligned + # with the next owned tile. + if cutlass.const_expr(self.has_bias): + bias_consumer_state.advance() + + # Pull the next tile from the CLC response slot. Both MMA + # warpgroups, the SFA cp.async group, and the scheduler warp + # participate as consumers; the parity flip below routes the tile + # to its owning warpgroup. + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + tile_parity = cutlass.Int32(1) - tile_parity + # End of for k_tile loop + # End of while loop + # End of MMA warp group + + # CLC scheduler warp: sole producer of the CLC tile stream. + elif warp_idx == self.sched_warp_id: + cute.arch.setmaxregister_decrease(self.sched_register_requirement) + clc_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_clc_stage + ) + + while work_tile.is_valid_tile: + clc_pipeline.producer_acquire(clc_producer_state) + tile_sched.advance_to_next_work( + clc_pipeline.producer_get_barrier(clc_producer_state) + ) + clc_producer_state.advance() + + # Self-consume to learn whether the next tile is valid. + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + clc_pipeline.producer_tail(clc_producer_state) + # Dedicated residual G2S load warp (only compiled when beta != 0). Walks + # the full tile stream as a CLC consumer -- like the TMA load warp -- and + # issues the residual im2col G2S load for every output tile, one per + # epi-subtile, into sResidual. Both math warpgroups consume their own + # tile's residual from the same pipeline. + elif warp_idx == self.residual_tma_warp_id: + # Dedicated residual G2S load warp. The whole body is compiled only + # when beta != 0; at beta == 0 this warp just sets the register floor + # like the other idle round-up warps (has_residual gates on the same + # const as the CLC consumer count and the smem/pipeline allocation, so + # nothing here touches the None residual handles). + cute.arch.setmaxregister_decrease(self.sched_register_requirement) + if cutlass.const_expr(self.has_residual): + residual_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.res_stage + ) + # Same gmem/smem TMA partition the epilogue consumer uses, so the + # producer stages subtiles in the exact grid the consumer reads. + sResidual_for_tma = cute.group_modes(sResidual, 0, 2) + while work_tile.is_valid_tile: + tile_coord_mnl = work_tile.tile_idx + gResidual_slice = gResidual_mnl[(None, None, *tile_coord_mnl)] + tcg_residual = cute.zipped_divide(gResidual_slice, self.epi_tile) + bGS_sResidual, bGS_gResidual = cpasync.tma_partition( + tma_atom_residual, + 0, + cute.make_layout(1), + sResidual_for_tma, + tcg_residual, + ) + residual_rest_m = bGS_gResidual.shape[1][0] + residual_rest_n = bGS_gResidual.shape[1][1] + for epi_m in cutlass.range_constexpr(residual_rest_m): + for epi_n in cutlass.range_constexpr(residual_rest_n): + residual_pipeline.producer_acquire(residual_producer_state) + cute.copy( + tma_atom_residual, + bGS_gResidual[(None, (epi_m, epi_n))], + bGS_sResidual[(None, residual_producer_state.index)], + tma_bar_ptr=residual_pipeline.producer_get_barrier( + residual_producer_state + ), + ) + residual_pipeline.producer_commit(residual_producer_state) + residual_producer_state.advance() + + # Walk the tile stream in lockstep with the epilogue consumers. + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + residual_pipeline.producer_tail(residual_producer_state) + # Unused warps rounded up by the warpgroup register-realloc requirement. + # Same warpgroup as the scheduler warp, so it must use the same count. + else: + cute.arch.setmaxregister_decrease(self.sched_register_requirement) + return + + def epilog_smem_copy_and_partition( + self, + tiled_mma: cute.TiledMma, + accumulators: cute.Tensor, + tidx: cutlass.Int32, + sD: cute.Tensor, + ) -> Tuple[ + cute.TiledCopy, + cute.TiledCopy, + cute.Tensor, + cute.Tensor, + cute.Tensor, + cute.Layout, + ]: + """Make the tiledCopy for the register-to-smem store and partition with it. + + :param tiled_mma: Tiled MMA object, whose C partition sets the thread-value + layout the store has to match + :type tiled_mma: cute.TiledMma + :param accumulators: The register accumulator tensor + :type accumulators: cute.Tensor + :param tidx: The thread index within the CTA + :type tidx: cutlass.Int32 + :param sD: The shared memory staging tensor for D + :type sD: cute.Tensor + + :return: A tuple containing (tiled_copy_r2s, thr_copy_r2s, tRS_sD, tRS_rAcc, + tRS_rD, tRS_rD_layout) where: + - tiled_copy_r2s: The tiled copy operation for register to smem copy(r2s) + - thr_copy_r2s: This thread's slice of it, which the residual read-back + reuses so its register fragment is congruent with tRS_rD + - tRS_sD: The partitioned tensor D (smem destination) + - tRS_rAcc: The accumulator retiled to the store's fragment + - tRS_rD: The register tensor the epilogue narrows into + - tRS_rD_layout: Its layout, which the bias and residual fragments reuse + :rtype: Tuple[cute.TiledCopy, cute.TiledCopy, cute.Tensor, cute.Tensor, + cute.Tensor, cute.Layout] + """ + copy_atom_r2s = sm120_utils.sm120_get_smem_store_op( + self.d_layout, + elem_ty_d=self.d_dtype, + elem_ty_acc=self.acc_dtype, + ) + + # StMatrix is a 16-bit instruction; pass Float16 so the partition + # geometry (which thread holds which element) is consistent + # regardless of d_dtype. The actual rmem->smem store is performed + # by copy_atom_r2s above, which is selected for the output type. + copy_atom_C = cute.make_copy_atom( + cute.nvgpu.warp.StMatrix8x8x16bOp( + self.d_layout.is_m_major_c(), + 2, + ), + cutlass.Float16, + ) + + tiled_copy_C_Atom = cute.make_tiled_copy_C_atom(copy_atom_C, tiled_mma) + + tiled_copy_r2s = cute.make_tiled_copy_S( + copy_atom_r2s, + tiled_copy_C_Atom, + ) + + thr_copy_r2s = tiled_copy_r2s.get_slice(tidx % 128) + # (R2S, R2S_M, R2S_N, PIPE_D) + tRS_sD = thr_copy_r2s.partition_D(sD) + # (R2S, R2S_M, R2S_N) + tRS_rAcc = tiled_copy_r2s.retile(accumulators) + + # Allocate D registers. + rD_shape = cute.shape(thr_copy_r2s.partition_S(sD)) + tRS_rD_layout = cute.make_layout(rD_shape[:3]) + tRS_rD = cute.make_rmem_tensor(tRS_rD_layout.shape, self.acc_dtype) + return tiled_copy_r2s, thr_copy_r2s, tRS_sD, tRS_rAcc, tRS_rD, tRS_rD_layout + + def epilog_gmem_copy_and_partition( + self, + tma_atom_d: cute.CopyAtom, + gD_mnl_slice: cute.Tensor, + sD: cute.Tensor, + ) -> Tuple[cute.Tensor, cute.Tensor]: + """Partition shared memory (source) and global memory (destination) for the + TMA store. + + :param tma_atom_d: The TMA copy atom for the D store + :type tma_atom_d: cute.CopyAtom + :param gD_mnl_slice: This CTA's tile of the global tensor D + :type gD_mnl_slice: cute.Tensor + :param sD: The shared memory staging tensor for D + :type sD: cute.Tensor + + :return: A tuple containing (bSG_sD, bSG_gD), the partitioned shared memory + and global tensors + :rtype: Tuple[cute.Tensor, cute.Tensor] + """ + sepi_for_tma_partition = cute.group_modes(sD, 0, 2) + tcgc_for_tma_partition = cute.zipped_divide(gD_mnl_slice, self.epi_tile) + + # ((ATOM_V, REST_V), EPI_M, EPI_N) + return cpasync.tma_partition( + tma_atom_d, + 0, + cute.make_layout(1), + sepi_for_tma_partition, + tcgc_for_tma_partition, + ) + + def _compute_stages( + self, + tile_shape_mnk: tuple[int, int, int], + a_dtype: type[cutlass.Numeric], + b_dtype: type[cutlass.Numeric], + sf_dtype: type[cutlass.Numeric], + epi_tile: tuple[int, int], + d_dtype: type[cutlass.Numeric], + smem_capacity: int, + occupancy: int, + has_residual: bool = False, + ) -> tuple[int, int]: + """Computes the number of A/B and epilogue stages. + + The A/B stage count is chosen so the full SharedStorage (every pipeline + mbar, whose count scales with the stage count, plus every operand array + with its 1024B alignment padding) fits in smem. The count is derived by + assembling the actual SharedStorage struct for a candidate stage count and + reading its exact size_in_bytes(), then taking the largest that fits -- + rather than dividing capacity by a per-stage byte estimate that omits the + alignment padding and the stage-scaled mbar bytes. + + :param tile_shape_mnk: The shape (M, N, K) of the CTA tile. + :param smem_capacity: Total available shared memory capacity in bytes. + :param occupancy: CTAs per SM. Always 1 for this persistent kernel. + + :return: (A/B operand stages, epilogue stages) + :rtype: tuple[int, int] + """ + epi_stage_max = (tile_shape_mnk[1] // epi_tile[1]) * ( + tile_shape_mnk[0] // epi_tile[0] + ) + epi_stage = min(epi_stage_max, 4) + res_stage = 1 + + buffer_align_bytes = self.buffer_align_bytes + num_clc_stage = self.num_clc_stage + has_bias = self.has_bias + cta_n = tile_shape_mnk[1] + + def smem_bytes(ab_stage: int) -> int: + """Exact SharedStorage byte size for a candidate stage count. + + Builds the operand/SF/epi/residual layouts for this ``ab_stage`` and + assembles the SharedStorage this kernel allocates. Every field mirrors + the runtime SharedStorage in __call__; keep the two in sync so the byte + total is exact. + """ + ( + a_smem, + b_smem, + sfa_smem, + sfb_smem, + epi_smem, + res_smem, + ) = self._make_smem_layouts( + tile_shape_mnk, + epi_tile, + a_dtype, + self.a_layout, + b_dtype, + self.b_layout, + ab_stage, + d_dtype, + self.d_layout, + epi_stage, + res_stage, + sf_vec_size=self.sf_vec_size, + tiled_mma=self.tiled_mma, + sf_tile_shape_mnk=self.sf_tile_shape_mnk, + ) + + @cute.struct + class ProbeStorage: + """Field-for-field mirror of the runtime SharedStorage, sized for + this stage count so size_in_bytes() gives the exact allocation.""" + + mainloop_pipeline_array_ptr: cute.struct.MemRange[ + cutlass.Int64, ab_stage * 2 + ] + math_wg_order_barrier_array_ptr: cute.struct.MemRange[cutlass.Int64, 4] + clc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, num_clc_stage * 2] + clc_response: cute.struct.Align[ + cute.struct.MemRange[cutlass.Int32, num_clc_stage * 4], 16 + ] + sA: cute.struct.Align[ + cute.struct.MemRange[a_dtype, cute.cosize(a_smem)], + buffer_align_bytes, + ] + sB: cute.struct.Align[ + cute.struct.MemRange[b_dtype, cute.cosize(b_smem)], + buffer_align_bytes, + ] + sSFA: cute.struct.Align[ + cute.struct.MemRange[sf_dtype, cute.cosize(sfa_smem)], + buffer_align_bytes, + ] + sSFB: cute.struct.Align[ + cute.struct.MemRange[sf_dtype, cute.cosize(sfb_smem)], + buffer_align_bytes, + ] + sD: cute.struct.Align[ + cute.struct.MemRange[d_dtype, cute.cosize(epi_smem)], + buffer_align_bytes, + ] + sBias: cute.struct.Align[ + cute.struct.MemRange[ + d_dtype, self.bias_stage * cta_n if has_bias else 0 + ], + buffer_align_bytes, + ] + sResidual: cute.struct.Align[ + cute.struct.MemRange[ + d_dtype, cute.cosize(res_smem) if has_residual else 0 + ], + buffer_align_bytes, + ] + residual_pipeline_array_ptr: cute.struct.MemRange[ + cutlass.Int64, res_stage * 2 if has_residual else 0 + ] + bias_pipeline_array_ptr: cute.struct.MemRange[ + cutlass.Int64, self.bias_stage * 2 if has_bias else 0 + ] + + return ProbeStorage.size_in_bytes() + + # A single A/B stage must fit: a tile whose one-stage SharedStorage already + # exceeds capacity is unimplementable and is rejected upstream by + # can_implement, so the lower-bound search below can assume >= 1 fits. + if smem_bytes(1) > smem_capacity: + raise RuntimeError( + "tile too large for even one A/B stage; can_implement should have " + "rejected it before reaching stage selection" + ) + + # Compute a lower bound that is guaranteed to fit, then scan upward once + # (never downward) to the largest A/B stage whose exact SharedStorage + # fits. Each added stage grows the four staged arrays (sA/sB/sSFA/sSFB) by + # their raw per-stage bytes; the Align[1024] rounding of each array is a + # one-time boundary crossing over the whole growth, so it contributes at + # most 4*1024 B total. Bounding the alignment loss as that constant in the + # numerator (rather than inflating every stage's slope) keeps the lower + # bound within a few stages of the true maximum while still guaranteeing + # cost(lb) <= cost(1) + (lb-1)*slope + 4*1024 <= capacity. size_in_bytes + # is monotonic in the stage count, so the upward scan lands on the maximum. + a_one, b_one, sfa_one, sfb_one, _, _ = self._make_smem_layouts( + tile_shape_mnk, + epi_tile, + a_dtype, + self.a_layout, + b_dtype, + self.b_layout, + 1, + d_dtype, + self.d_layout, + epi_stage, + res_stage, + sf_vec_size=self.sf_vec_size, + tiled_mma=self.tiled_mma, + sf_tile_shape_mnk=self.sf_tile_shape_mnk, + ) + slope = ( + cute.cosize(a_one) * a_dtype.width // 8 + + cute.cosize(b_one) * b_dtype.width // 8 + + cute.cosize(sfa_one) * sf_dtype.width // 8 + + cute.cosize(sfb_one) * sf_dtype.width // 8 + + 2 * 8 # two mainloop pipeline mbar (Int64) per stage + ) + align_loss = 4 * buffer_align_bytes + ab_stage = max(1, 1 + (smem_capacity - smem_bytes(1) - align_loss) // slope) + while smem_bytes(ab_stage + 1) <= smem_capacity: + ab_stage += 1 + return ab_stage, epi_stage + + @staticmethod + def _make_smem_layouts( + tile_shape_mnk: tuple[int, int, int], + epi_tile: tuple[int, int], + a_dtype: type[cutlass.Numeric], + a_layout: cute.Layout, + b_dtype: type[cutlass.Numeric], + b_layout: cute.Layout, + ab_stage: int, + d_dtype: type[cutlass.Numeric], + d_layout: cute.Layout, + epi_stage: int, + res_stage: int, + sf_vec_size: int, + tiled_mma: cute.TiledMma, + sf_tile_shape_mnk: tuple[int, int, int] = None, + ) -> tuple[cute.ComposedLayout, cute.ComposedLayout, cute.ComposedLayout]: + """Create shared memory layouts for the A, B and D tensors. + + :param tile_shape_mnk: CTA tile shape (M,N,K). Sizes A/B/D/residual. + :type tile_shape_mnk: Tuple[int, int, int] + :param sf_tile_shape_mnk: CTA tile shape whose K sizes the SF SMEM. For + MXFP8 the SF K tile (128) decouples from the A/B K tile (64); defaults + to tile_shape_mnk when they match (NVFP4). + :param epi_tile: Epilogue tile shape + :type epi_tile: Tuple[int, int] + :param a_dtype: Data type for matrix A + :type a_dtype: type[cutlass.Numeric] + :param a_layout: Layout for matrix A + :type a_layout: Layout + :param b_dtype: Data type for matrix B + :type b_dtype: type[cutlass.Numeric] + :param b_layout: Layout for matrix B + :type b_layout: Layout + :param ab_stage: Number of stages for A/B tensors + :type ab_stage: int + :param d_dtype: Data type for output matrix D + :type d_dtype: type[cutlass.Numeric] + :param d_layout: leading dimension of the output matrix D + :type d_layout: Layout + :param epi_stage: Number of epilogue stages + :type epi_stage: int + + :return: Tuple of shared memory layouts for A, B and the epilogue + :rtype: Tuple[cute.ComposedLayout, cute.ComposedLayout, cute.ComposedLayout] + """ + if sf_tile_shape_mnk is None: + sf_tile_shape_mnk = tile_shape_mnk + + a_smem_shape = cute.slice_(tile_shape_mnk, (None, 0, None)) + + a_is_k_major = a_layout.is_k_major_a() + b_is_k_major = b_layout.is_k_major_b() + a_major_mode_size = tile_shape_mnk[2 if a_is_k_major else 0] + + a_smem_layout_atom = cute.nvgpu.warpgroup.make_smem_layout_atom( + sm90_utils.get_smem_layout_atom( + a_layout, + a_dtype, + a_major_mode_size, + ), + a_dtype, + ) + a_smem_layout_staged = cute.tile_to_shape( + a_smem_layout_atom, + cute.append(a_smem_shape, ab_stage), + order=(0, 1, 2) if a_is_k_major else (1, 0, 2), + ) + + b_smem_shape = cute.slice_(tile_shape_mnk, (0, None, None)) + + b_major_mode_size = tile_shape_mnk[2 if b_is_k_major else 1] + b_smem_layout_atom = cute.nvgpu.warpgroup.make_smem_layout_atom( + sm90_utils.get_smem_layout_atom( + b_layout, + b_dtype, + b_major_mode_size, + ), + b_dtype, + ) + b_smem_layout_staged = cute.tile_to_shape( + b_smem_layout_atom, + cute.append(b_smem_shape, ab_stage), + order=(0, 1, 2) if b_is_k_major else (1, 0, 2), + ) + + sfa_smem_layout_staged = blockscaled_utils.sm120_make_smem_layout_sfa( + tiled_mma, + sf_tile_shape_mnk, + sf_vec_size, + ab_stage, + ) + + sfb_smem_layout_staged = _sm120_make_smem_layout_sfb( + tiled_mma, + sf_tile_shape_mnk, + sf_vec_size, + ab_stage, + ) + + d_smem_shape = epi_tile + d_major_mode_size = epi_tile[1] if d_layout.is_n_major_c() else epi_tile[0] + d_smem_layout_atom = cute.nvgpu.warpgroup.make_smem_layout_atom( + sm90_utils.get_smem_layout_atom( + d_layout, + d_dtype, + d_major_mode_size, + ), + d_dtype, + ) + epi_smem_layout_staged = cute.tile_to_shape( + d_smem_layout_atom, + cute.append(d_smem_shape, epi_stage), + order=(1, 0, 2) if d_layout.is_m_major_c() else (0, 1, 2), + ) + # Residual staging layout: same per-subtile atom as the D store, sized by + # res_stage. Single-buffered (res_stage=1) so it costs minimal smem and + # leaves the A/B pipeline deep. + res_smem_layout_staged = cute.tile_to_shape( + d_smem_layout_atom, + cute.append(d_smem_shape, res_stage), + order=(1, 0, 2) if d_layout.is_m_major_c() else (0, 1, 2), + ) + + return ( + a_smem_layout_staged, + b_smem_layout_staged, + sfa_smem_layout_staged, + sfb_smem_layout_staged, + epi_smem_layout_staged, + res_smem_layout_staged, + ) + + @staticmethod + def can_implement( + tile_shape_mnk: Tuple[int, int, int], + ab_dtype: Type[cutlass.Numeric], + sf_dtype: Type[cutlass.Numeric], + acc_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + c: int, + k: int, + filter_trs: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + dil_dhw: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], + ) -> None: + """Rejects configurations this kernel cannot compile or would miscompute. + + Every constraint here is decidable from the tile shape, the input format + and the channel counts, so it is checked before any compilation rather + than surfacing as a deep layout error or a wrong result. + + :param tile_shape_mnk: CTA tile shape (M, N, K) + :param ab_dtype: A/B element type + :param sf_dtype: Scale-factor element type + :param acc_dtype: Accumulator element type + :param d_dtype: Output element type + :param sf_vec_size: Channels one scale factor covers + :param c: Input channel count + :param k: Output channel count, which sets how many N tiles are needed + :param filter_trs: Filter extents (T, R, S) + :param stride_dhw: Convolution stride per spatial dimension + :param dil_dhw: Dilation per spatial dimension + :param upper_padding_dhw: Trailing padding per spatial dimension + :param lower_padding_dhw: Leading padding per spatial dimension + + :raises testing.CantImplementError: If the configuration is unsupported + """ + # Two validated input formats, each pinned to its scale-factor block size: + # NVFP4 : Float4E2M1FN A/B, sf_vec_size=16 (E4M3 scale). + # MXFP8 : Float8E4M3FN A/B, sf_vec_size=32 (E8M0 scale). + is_mxfp8 = ab_dtype is cutlass.Float8E4M3FN + if ab_dtype is cutlass.Float4E2M1FN: + if sf_dtype is not cutlass.Float8E4M3FN or sf_vec_size != 16: + raise testing.CantImplementError( + "Float4E2M1FN A/B requires sf_dtype=Float8E4M3FN and " + f"sf_vec_size=16, got sf_dtype={sf_dtype}, " + f"sf_vec_size={sf_vec_size}" + ) + elif is_mxfp8: + if sf_dtype is not cutlass.Float8E8M0FNU or sf_vec_size != 32: + raise testing.CantImplementError( + "Float8E4M3FN A/B (MXFP8) requires sf_dtype=Float8E8M0FNU and " + f"sf_vec_size=32, got sf_dtype={sf_dtype}, " + f"sf_vec_size={sf_vec_size}" + ) + else: + raise testing.CantImplementError( + f"This path supports Float4E2M1FN (NVFP4) and Float8E4M3FN (MXFP8) " + f"A/B, got ab_dtype={ab_dtype}" + ) + # tile_m must be exactly 128: the block-scaled SF SMEM layout tiles the M + # extent in units of blk_mn=128 (one BlockScaledBasicChunk row), so a tile_m + # that is not a whole number of chunks fails that layout's own assert. + if tile_shape_mnk[0] != 128: + raise testing.CantImplementError( + f"tile_m must be 128 (the SF chunk's row count), got " + f"{tile_shape_mnk[0]}" + ) + # tile_k for NVFP4 is 64 or 128: the MMA fills K in 64-channel steps and the + # SF SMEM layout needs a whole 4-block scale-factor atom (also 64 channels) + # per K tile, so 64 is the floor and nothing between is expressible. Prefer + # the value that equals C. 128 is the widest value tested. + # + # MXFP8 uses a 64 K tile so the fp8 A/B, at twice fp4's bytes per stage, + # keeps a deep A/B pipeline in the smaller SM120 SMEM; its SF K tile stays + # at one whole chunk of 4*vec32 = 128 channels, so one SF chunk serves two + # A/B tiles. + allowed_tile_k = (64,) if is_mxfp8 else (64, 128) + allowed_tile_n = (64, 96, 128) + if tile_shape_mnk[2] not in allowed_tile_k: + raise testing.CantImplementError( + f"tile_k must be one of {allowed_tile_k} for {ab_dtype}, got " + f"{tile_shape_mnk[2]}" + ) + # tile_n is bounded by the scale-factor chunk geometry. A chunk covers 128 + # output channels and is indivisible, and the SF layout a tile reads assumes + # the tile starts on a chunk boundary. Values outside this set are each + # ruled out for their own reason: a tile_n that is not a multiple of 32 + # misses the SF block granularity and would gain nothing anyway (the MMA + # fills N in 32-wide units and rounds up); 32 has no ldmatrix partition for + # the scale factors, which reach the MMA through ldmatrix; above 128 the N + # basic block still describes only one chunk's four 32-row units, so the + # layout stops covering the tile. + tile_n = tile_shape_mnk[1] + if tile_n not in allowed_tile_n: + raise testing.CantImplementError( + f"tile_n must be one of {allowed_tile_n} for {ab_dtype} (bounded by " + f"the 128-channel SF chunk and its 32-row blocks), got {tile_n}" + ) + # 64 and 128 divide a chunk, so every N tile sits inside one. 96 does not, + # so only the first tile is chunk-aligned; it is restricted to problems that + # need a single N tile, which is where it wins anyway by filling N exactly. + n_tiles = (k + tile_n - 1) // tile_n + if 128 % tile_n != 0 and n_tiles > 1: + raise testing.CantImplementError( + f"tile_n={tile_n} does not divide the 128-channel SF chunk, so only " + f"the first N tile starts on a chunk boundary; it needs a single N " + f"tile but K={k} needs {n_tiles}. Use tile_n=64 or 128." + ) + # C neither has to fill the K tile nor be a whole number of scale-factor + # atoms. A partial trailing K tile reads zero-filled A and B channels from + # the TMA, whose global channel extent is C; SFA re-reads a chunk that is + # already in bounds because the host rounds each pixel's row up to a whole + # atom and zero-fills the tail; and SFB is addressed per filter position at + # a span padded to whole K tiles. So the tail contributes nothing. + # + # One bound survives: the 16-byte alignment the TMA needs on a contiguous + # axis. C is the contiguous axis of A and B, so a pixel's row -- C elements + # of ab_dtype -- is the stride of the axis next to it, and the descriptor + # requires that stride to be a multiple of 16 bytes (C % 32 == 0 for FP4, + # C % 16 == 0 for FP8). K is the contiguous axis of D and carries the same + # requirement against the output element type. + _check_tensor_alignment(c, k, ab_dtype, d_dtype) + # A tighter bound on C than that alignment: it must be a whole number of + # scale-factor atoms, 4 * sf_vec_size channels. SFA's 4-byte cp.async carries + # exactly one atom and walks a row stride of C / sf_vec_size factors, which + # this keeps a multiple of the 4 the transfer needs. A C that splits an atom + # would leave that stride unaligned, and padding the host allocation to hide + # it would put the factors a position owns at an offset the kernel does not + # address. + sf_atom_channels = sf_vec_size * 4 + if c % sf_atom_channels != 0: + raise testing.CantImplementError( + f"{ab_dtype} requires C to be a multiple of {sf_atom_channels}, " + f"got C = {c}." + ) + # The block-scaled MMA accumulates in FP32 only -- its op rejects any other + # accumulator type when it is built, past the point a host check can name the + # problem. + if acc_dtype is not cutlass.Float32: + raise testing.CantImplementError( + f"acc_dtype must be Float32, got {acc_dtype}" + ) + # D is a 16-bit float. The epilogue converts the accumulator straight to D, + # and a sub-byte output would additionally need its own scale-factor tensor + # quantized from that accumulator, which nothing here emits. + allowed_d_dtype = {cutlass.BFloat16, cutlass.Float16} + if d_dtype not in allowed_d_dtype: + raise testing.CantImplementError( + f"D must be BFloat16 or Float16, got d_dtype={d_dtype}" + ) + _check_im2col_descriptor_limits( + filter_trs, stride_dhw, dil_dhw, upper_padding_dhw, lower_padding_dhw + ) + + +@cute.jit +def cvt_sf_MKL_to_M32x4xrm_K4xrk_L( + sf_ref_tensor: cute.Tensor, + sf_mma_tensor: cute.Tensor, +): + """Convert scale factor tensor from MKL layout to mma specification M(32x4xrest_m)xK(4xrest_k)xL layout""" + # sf_mma_tensor has flatten shape (32, 4, rest_m, 4, rest_k, l) + # group to ((32, 4, rest_m), (4, rest_k), l) + sf_mma_tensor = cute.group_modes(sf_mma_tensor, 0, 3) + sf_mma_tensor = cute.group_modes(sf_mma_tensor, 1, 3) + for i in cutlass.range(cute.size(sf_ref_tensor)): + mkl_coord = sf_ref_tensor.layout.get_hier_coord(i) + sf_mma_tensor[mkl_coord] = sf_ref_tensor[mkl_coord] + + +def _sm120_make_smem_layout_sfb( + tiled_mma: cute.TiledMma, + tile_shape_mnk: cute.Tile, + sf_vec_size: int, + num_stages: int, +) -> cute.Layout: + """ + Make smem layout for SFB based on: + 1. BlockScaledBasicChunk + 2. MMA tiler shape + 3. Scale factor vector size + 4. Number of stages + + :param tiled_mma: The tiled MMA + :type tiled_mma: cute.TiledMma + :param tile_shape_mnk: The mma tiler shape + :type tile_shape_mnk: cute.Tile + :param sf_vec_size: The scale factor vector size + :type sf_vec_size: int + :param num_stages: The number of stages + :type num_stages: int + + :return: Smem layout for SFB + :rtype: cute.Layout + """ + + # A single indivisible block will hold 4 scale factors of 128 rows/columns (A/B matrix). + # 4 is chosen to make consecutive 32bits of data to have scale factors for only a single row(col). + blk_mn = 128 + blk_sf = 4 + blk_elems = blk_mn * blk_sf + + assert sf_vec_size == 16 or sf_vec_size == 32, "sf_vec_size must be 16 or 32" + assert isinstance(tile_shape_mnk, tuple) + + # Below a chunk the layout describes that chunk's leading tile_n columns, so + # tile_n only has to land on the 32-row basic block. At or above a chunk the N + # mode still spans exactly blk_sf blocks, so a tile_n that is not a whole number + # of chunks would leave the layout covering fewer columns than the tile. + assert tile_shape_mnk[1] % 32 == 0 and ( + tile_shape_mnk[1] < blk_mn or tile_shape_mnk[1] % blk_mn == 0 + ), ( + "tile_shape_mnk[1] must be a multiple of the 32-row SF basic block, and " + f"either below blk_mn={blk_mn} or a whole number of chunks" + ) + + assert tile_shape_mnk[2] % sf_vec_size == 0, ( + "tile_shape_mnk[2] must be divisible by sf_vec_size" + ) + + mma_nsf = tiled_mma.shape_mnk[2] // sf_vec_size + + # A chunk holds blk_mn=128 columns and is indivisible in SMEM, so a tile_n + # narrower than a chunk describes the leading tile_n columns of one chunk: + # the N atom shrinks to (32, tile_n/32) while every stride, including the + # K-mode chunk stride, still steps over whole chunks. + n_chunks = max(1, tile_shape_mnk[1] // blk_mn) + n_atom_blocks = min(blk_sf, tile_shape_mnk[1] // 32) + + mn_basic_block_shape = (32, n_atom_blocks) + mn_basic_block_stride = (16, 4) + k_basic_block_shape = (sf_vec_size, mma_nsf) + k_basic_block_stride = (0, 1) + + sSFA_shapeN = (mn_basic_block_shape, n_chunks) + sSF_strideN = (mn_basic_block_stride, blk_elems) + + assert tile_shape_mnk[2] % (blk_sf * mma_nsf) == 0, ( + "tile_shape_mnk[2] must be divisible by blk_sf * mma_nsf" + ) + + sSFA_shapeK = ( + k_basic_block_shape, + blk_sf // mma_nsf, + tile_shape_mnk[2] // sf_vec_size // blk_sf, + ) + sSF_strideK = ( + k_basic_block_stride, + mma_nsf, + n_chunks * blk_elems, + ) + + sSFA_shape = (sSFA_shapeN, sSFA_shapeK) + sSFA_stride = (sSF_strideN, sSF_strideK) + + smem_layout = cute.make_layout(sSFA_shape, stride=sSFA_stride) + + # A stage always spans whole chunks, so its stride comes from the chunk + # count rather than the layout's cosize: a tile_n narrower than a chunk + # addresses only part of the chunk it sits in, and a cosize would place the + # next stage inside that same chunk. + k_chunks = tile_shape_mnk[2] // sf_vec_size // blk_sf + stage_stride = n_chunks * blk_elems * k_chunks + + # (((Atom_Inst_M, Rest_M),(Atom_Inst_K, Rest_K)), MMA_M, MMA_K, STAGE) + sfb_smem_layout_staged = cute.append( + smem_layout, + cute.make_layout(num_stages, stride=stage_stride), + ) + + return sfb_smem_layout_staged + + +def sf_k_tile_channels(cta_tile_k: int, sf_vec_size: int) -> int: + """Channels one SF K tile spans. + + The SF SMEM atom holds 4 scale-factor blocks along K, so it covers + 4 * sf_vec_size channels and cannot be staged in part. A narrower A/B K tile + therefore still stages a whole atom, and several A/B tiles then share one SF + tile. Both the kernel and the host derive the SF cadence from here so they + cannot drift apart. + """ + return max(cta_tile_k, 4 * sf_vec_size) + + +def sfb_per_position_channels(c: int, cta_tile_k: int, sf_vec_size: int) -> int: + """Channels of SFB one filter position spans: C rounded up to whole SF K tiles. + + SFB runs channels and filter positions together in one flat K mode, so a K tile + that does not divide C would straddle two positions; rounding each position out + to whole SF K tiles puts every position on a tile boundary, which keeps a tile + inside one position and lets the consumer address it by its flat index. The + bound is the SF tile, the width of the TMA box: a C that a 64-channel A/B tile + divides still straddles two positions once a 128-channel SF chunk serves two of + those tiles. + + The padded channels pair with the B channels past C, which the TMA zero-fills, + so their scale factors never reach the result. + + The kernel derives the same span from its runtime channel count, so the padding + never reaches the compiled code. + """ + span = sf_k_tile_channels(cta_tile_k, sf_vec_size) + return -(-c // span) * span + + +def create_scale_factor_tensor_swizzled( + l: int, + mn: int, + k: int, + sf_vec_size: int, + dtype: Type[cutlass.Numeric], +) -> Tuple[torch.Tensor, cute.Tensor, torch.Tensor]: + """SFB gmem tensor in the swizzled BlockScaledBasicChunk MMA layout (TMA + load). Returns (f32 reference in MKL layout, cute tensor, backing storage).""" + + def ceil_div(a, b): + return (a + b - 1) // b + + sf_k = ceil_div(k, sf_vec_size) + ref_shape = (l, mn, sf_k) + atom_m = (32, 4) + atom_k = 4 + mma_shape = ( + l, + ceil_div(mn, atom_m[0] * atom_m[1]), + ceil_div(sf_k, atom_k), + atom_m[0], + atom_m[1], + atom_k, + ) + + ref_permute_order = (1, 2, 0) + mma_permute_order = (3, 4, 1, 5, 2, 0) + + ref_f32_cpu = cutlass_torch.create_and_permute_torch_tensor( + ref_shape, + torch.float32, + permute_order=ref_permute_order, + init_type=cutlass_torch.TensorInitType.RANDOM, + init_config=cutlass_torch.RandomInitConfig(min_val=1, max_val=3), + ) + cute_f32_cpu = cutlass_torch.create_and_permute_torch_tensor( + mma_shape, + torch.float32, + permute_order=mma_permute_order, + init_type=cutlass_torch.TensorInitType.RANDOM, + init_config=cutlass_torch.RandomInitConfig(min_val=0, max_val=1), + ) + + cvt_sf_MKL_to_M32x4xrm_K4xrk_L( + from_dlpack(ref_f32_cpu), + from_dlpack(cute_f32_cpu), + ) + + ref_f32_cpu = ( + ref_f32_cpu.permute(2, 0, 1) + .unsqueeze(-1) + .expand(l, mn, sf_k, sf_vec_size) + .reshape(l, mn, sf_k * sf_vec_size) + .permute(*ref_permute_order) + ) + ref_f32_cpu = ref_f32_cpu[:, :k, :] + # Round-trip the reference scale factors through the storage dtype so the + # reference dequant uses the exact values the kernel reads. E8M0 (MXFP8) + # snaps every scale to a power of two, so an un-rounded f32 reference would + # disagree with the kernel on every non-pow2 block. + ref_f32_cpu = ref_f32_cpu.to(cutlass_torch.dtype(dtype)).to(torch.float32) + + cute_f32 = cute_f32_cpu.cuda() + + cute_tensor, torch_tensor = cutlass_torch.cute_tensor_like( + cute_f32_cpu, + dtype, + is_dynamic_layout=True, + assumed_align=16, + ) + cute_tensor = cutlass_torch.convert_cute_tensor( + cute_f32, + cute_tensor, + dtype, + is_dynamic_layout=True, + ) + return ref_f32_cpu, cute_tensor, torch_tensor + + +def create_scale_factor_tensor_unswizzled( + l: int, + mn: int, + k: int, + sf_vec_size: int, + dtype: Type[cutlass.Numeric], +) -> Tuple[torch.Tensor, cute.Tensor, torch.Tensor]: + """SFA gmem tensor in natural (M, ceil(C/vec), L) unswizzled layout consumed + by the hand-written cp.async producer (the gather-XOR-swizzle theorem + forbids a TMA for SFA). Returns (f32 reference in MKL layout, cute tensor, + backing storage).""" + + def ceil_div(a, b): + return (a + b - 1) // b + + # The M-direction stride is sf_k and the 4-byte cp.async that walks it needs a + # multiple of 4 factors. k is C here, which can_implement already requires to be + # a whole number of scale-factor atoms, so the allocation carries no tail: the + # caller's buffer is exactly the factors the problem owns. + sf_k = ceil_div(k, sf_vec_size) + + sf_raw = torch.randint(1, 3, (mn, sf_k, l), dtype=torch.uint8, device="cpu") + torch_tensor = sf_raw.to(dtype=cutlass_torch.dtype(dtype)).cuda() + + cute_tensor = from_dlpack(torch_tensor, assumed_align=16) + cute_tensor.element_type = dtype + # Mark the SFA gmem tensor layout dynamic (leading dim = the C/SF axis) so + # its stride lowers to runtime SSA. Without it the static stride is baked + # into the cubin, tying the compiled kernel to one C (the device reads mSFA + # with this stride). + cute_tensor = cute_tensor.mark_layout_dynamic(leading_dim=1) + + ref_f32_cpu = ( + sf_raw[:, :sf_k, :] + .float() + .permute(2, 0, 1) + .unsqueeze(-1) + .expand(l, mn, sf_k, sf_vec_size) + .reshape(l, mn, sf_k * sf_vec_size) + .permute(1, 2, 0) + ) + ref_f32_cpu = ref_f32_cpu[:, :k, :] + + return ref_f32_cpu, cute_tensor, torch_tensor + + +# Compile-time epilogue activations, folded into the cubin as a Constexpr op +# (one activation per cubin). Each entry pairs the device-side op applied to the +# output fragment with the torch op used to build the reference. +EPILOGUE_ACTIVATIONS = { + "identity": { + "device": lambda x: x, + "ref": lambda x: x, + }, + "relu": { + "device": lambda x: cute.where(x > 0, x, cute.full_like(x, 0)), + "ref": torch.nn.functional.relu, + }, +} + + +def compile_conv( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + zpq: Tuple[int, int, int], + input_: cute.Tensor, + filter_: cute.Tensor, + output_: cute.Tensor, + sfa_: cute.Tensor, + sfb_: cute.Tensor, + bias_: Optional[cute.Tensor], + residual_: Optional[cute.Tensor], + alpha_: cutlass.Float32, + beta: float, + acc_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + tile_shape_mnk: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + upper_pad_dhw: Tuple[int, int, int], + lower_pad_dhw: Tuple[int, int, int], + dil_dhw: Tuple[int, int, int], + epilogue_op: cutlass.Constexpr = lambda x: x, + allow_hardware_query_failure: bool = False, +): + """Build the kernel object, resolve host launch config, and cute.compile it. + + Returns the compiled callable. Pad/stride/dilation and filter T/R/S are boxed + as cutlass.Int32 so the compiled entry scalars lower to runtime SSA (one cubin + serves any of those values); the caller re-boxes them at launch time. The stream a + compilation is handed never runs anything, so it is a fake one; the caller passes + the real stream when it launches. + """ + from cutlass.cute.runtime import make_fake_stream + + _, t, r, s = ktrs + + conv = Sm120BlockScaledPersistentDenseImplicitGemmFpropKernel( + acc_dtype, + sf_vec_size, + tile_shape_mnk, + ) + hardware_info = cutlass.utils.HardwareInfo() + try: + max_active_clusters = hardware_info.get_max_active_clusters(1) + except Exception: + if not allow_hardware_query_failure: + raise + max_active_clusters = 1 + stream = make_fake_stream() + + compile_opts_parts = [] + dump_dir = os.environ.get("CUTE_DSL_DUMP_DIR") + if dump_dir: + compile_opts_parts.extend( + [f"--dump-dir={dump_dir}", "--keep-cubin", "--keep-ptx"] + ) + ptx_version = os.environ.get("CUTE_DSL_PTX_VERSION") + if ptx_version: + if ptx_version.isdigit(): + ptx_version = f"+ptx{ptx_version}" + compile_opts_parts.append(f"--ptx-version={ptx_version}") + compile_opts = " ".join(compile_opts_parts) + + rt_conv_params = ( + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + cutlass.Int32(t), + cutlass.Int32(r), + cutlass.Int32(s), + ) + + return cute.compile( + conv, + input_, + filter_, + sfa_, + sfb_, + output_, + bias_, + residual_, + alpha_, + beta, + *rt_conv_params, + max_active_clusters, + stream, + epilogue_op, + no_cache=True, + **({"options": compile_opts} if compile_opts else {}), + ) + + +def run_conv( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + ab_dtype: Type[cutlass.Numeric], + sf_dtype: Type[cutlass.Numeric], + sf_vec_size: int, + d_dtype: Type[cutlass.Numeric], + acc_dtype: Type[cutlass.Numeric], + tile_shape_mnk: Tuple[int, int, int], + tolerance: float, + warmup_iterations: int, + iterations: int, + skip_ref_check: bool, + use_cold_l2: bool = False, + stride_dhw: Tuple[int, int, int] = (1, 1, 1), + upper_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + lower_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + dil_dhw: Tuple[int, int, int] = (1, 1, 1), + activation: str = "identity", + use_bias: bool = False, + alpha: float = 1.0, + beta: float = 0.0, +): + if not torch.cuda.is_available(): + raise RuntimeError("GPU is required to run this example") + + n, c, d, h, w = ncdhw + k, t, r, s = ktrs + # Reject malformed geometry before compute_zpq(), which floor-divides by each + # stride: a zero stride would raise ZeroDivisionError, and a negative stride or + # dilation would silently feed invalid extents into the host reference and the + # device address reconstruction. Turn both into a clear ValueError up front. + for name, values in { + "stride_dhw": stride_dhw, + "upper_pad_dhw": upper_pad_dhw, + "lower_pad_dhw": lower_pad_dhw, + "dil_dhw": dil_dhw, + }.items(): + if len(values) != 3: + raise ValueError(f"{name} must contain exactly 3 values, got {values}") + if any(v <= 0 for v in stride_dhw): + raise ValueError( + f"stride_dhw must contain only positive values, got {stride_dhw}" + ) + if any(v <= 0 for v in dil_dhw): + raise ValueError(f"dil_dhw must contain only positive values, got {dil_dhw}") + if t <= 0 or r <= 0 or s <= 0: + raise ValueError(f"TRS must be positive, got {(t, r, s)}") + Sm120BlockScaledPersistentDenseImplicitGemmFpropKernel.can_implement( + tile_shape_mnk, + ab_dtype, + sf_dtype, + acc_dtype, + d_dtype, + sf_vec_size, + c, + k, + (t, r, s), + stride_dhw, + dil_dhw, + upper_pad_dhw, + lower_pad_dhw, + ) + + z, p, q = compute_zpq( + (d, h, w), (t, r, s), stride_dhw, upper_pad_dhw, lower_pad_dhw, dil_dhw + ) + if z <= 0 or p <= 0 or q <= 0: + raise ValueError(f"Invalid output spatial shape: {(z, p, q)}") + + gemm_m = n * z * p * q + gemm_n = k + gemm_k = c * t * r * s + print("Running Blackwell Geforce SM120 NVFP4 fprop with:") + print(f"ncdhw: {ncdhw}, ktrs: {ktrs}, zpq: {(z, p, q)}") + print(f"implicit GEMM mnk: {(gemm_m, gemm_n, gemm_k)}") + print( + f"A/B dtype: {ab_dtype}, SF dtype: {sf_dtype}, D dtype: {d_dtype}, Acc dtype: {acc_dtype}" + ) + print(f"Tile Shape: {tile_shape_mnk}") + print(f"Skip reference checking: {skip_ref_check}") + + # Resolve the compile-time epilogue activation. The device op is folded into + # the cubin as a Constexpr; the reference op mirrors it on the host. + if activation not in EPILOGUE_ACTIVATIONS: + raise ValueError( + f"Unsupported activation {activation!r}; " + f"choose from {sorted(EPILOGUE_ACTIVATIONS)}" + ) + epilogue_op = EPILOGUE_ACTIVATIONS[activation]["device"] + + input_tensor, filter_tensor, output_tensor = prepare_tensors(ncdhw, ktrs, (z, p, q)) + input_, input_storage = create_cute_tensor(input_tensor, ab_dtype, leading_dim=4) + filter_, filter_storage = create_cute_tensor(filter_tensor, ab_dtype, leading_dim=4) + output_, output_storage = create_cute_tensor(output_tensor, d_dtype, leading_dim=4) + + sfa_ref, sfa_tensor, sfa_storage = create_scale_factor_tensor_unswizzled( + 1, n * d * h * w, c, sf_vec_size, sf_dtype + ) + # SFB is allocated with every filter position padded out to a whole number of SF + # K tiles, the same span the kernel derives at runtime, so every position starts + # on a tile boundary and no K tile addresses past the one it belongs to. The + # padded channels pair with the B channels past C, which the TMA zero-fills, so + # their scale factors never reach the result and only have to be finite. + sfb_c_span = sfb_per_position_channels(c, tile_shape_mnk[2], sf_vec_size) + sfb_ref, sfb_tensor, sfb_storage = create_scale_factor_tensor_swizzled( + 1, gemm_n, sfb_c_span * t * r * s, sf_vec_size, sf_dtype + ) + + # Per-output-channel bias (length gemm_n = K), dtype matches the output. + # Added to the accumulator in FP32 as D = act(alpha*acc + bias), broadcast + # across every spatial output position. + if use_bias: + bias_storage = torch.randn(gemm_n, dtype=torch.float32, device="cuda").to( + cutlass_torch.dtype(d_dtype) + ) + bias_ = from_dlpack(bias_storage, assumed_align=16) + else: + bias_storage, bias_ = None, None + + # Per-tensor alpha as an FP32 runtime scalar (D = act(alpha*acc + bias + + # beta*residual)). Boxed once and passed at both cute.compile and every + # launch so one cubin serves any alpha value. + alpha_ = cutlass.Float32(alpha) + + # Residual (C): a per-element (N,Z,P,Q,K) tensor sharing the output's + # shape/layout/dtype, added in FP32 as beta*residual. beta is a compile-time + # constant, so beta == 0 needs no tensor and compiles the path away. + has_residual = beta != 0.0 + if has_residual: + residual_f32 = torch.randn((n, z, p, q, k), dtype=torch.float32, device="cuda") + residual_, residual_storage = create_cute_tensor( + residual_f32, d_dtype, leading_dim=4 + ) + else: + residual_, residual_storage, residual_f32 = None, None, None + + torch_stream = torch.cuda.Stream() + current_stream = cuda.CUstream(torch_stream.cuda_stream) + # Box pad/stride/dilation and filter T/R/S as cutlass.Int32 so the launch + # scalars lower to runtime SSA: one cubin serves any pad/stride/dilation/ + # T/R/S. Order must match the __call__ signature (compile_conv boxes the + # same values independently for the cute.compile entry). + rt_conv_params = ( + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + cutlass.Int32(t), + cutlass.Int32(r), + cutlass.Int32(s), + ) + + compiled_conv = compile_conv( + ncdhw, + ktrs, + (z, p, q), + input_, + filter_, + output_, + sfa_tensor, + sfb_tensor, + bias_, + residual_, + alpha_, + beta, + acc_dtype, + sf_vec_size, + tile_shape_mnk, + stride_dhw, + upper_pad_dhw, + lower_pad_dhw, + dil_dhw, + epilogue_op=epilogue_op, + allow_hardware_query_failure=skip_ref_check, + ) + print("Compiled conv kernel.", flush=True) + + # The inputs initialize on other streams; drain them so the kernel's + # stream never reads a tensor mid-initialization. + torch.cuda.synchronize() + compiled_conv( + input_, + filter_, + sfa_tensor, + sfb_tensor, + output_, + bias_, + residual_, + alpha_, + *rt_conv_params, + current_stream, + ) + torch_stream.synchronize() + + if not skip_ref_check: + sfa_expanded = sfa_ref.squeeze(-1).reshape(n, d, h, w, c).cuda() + sfb_expanded = ( + sfb_ref.squeeze(-1).reshape(k, t, r, s, sfb_c_span)[..., :c].cuda() + ) + scaled_input = input_tensor.float() * sfa_expanded + scaled_filter = filter_tensor.float() * sfb_expanded + scaled_input_ncdhw = scaled_input.permute(0, 4, 1, 2, 3).contiguous() + scaled_filter_kctrs = scaled_filter.permute(0, 4, 1, 2, 3).contiguous() + if upper_pad_dhw == lower_pad_dhw: + ref_nkzpq = F.conv3d( + scaled_input_ncdhw, + scaled_filter_kctrs, + padding=upper_pad_dhw, + stride=stride_dhw, + dilation=dil_dhw, + ) + else: + padded = F.pad( + scaled_input_ncdhw, + ( + lower_pad_dhw[2], + upper_pad_dhw[2], + lower_pad_dhw[1], + upper_pad_dhw[1], + lower_pad_dhw[0], + upper_pad_dhw[0], + ), + ) + ref_nkzpq = F.conv3d( + padded, + scaled_filter_kctrs, + stride=stride_dhw, + dilation=dil_dhw, + ) + ref = ref_nkzpq.permute(0, 2, 3, 4, 1).contiguous() + # Scale the accumulator by alpha in FP32, then add the per-output-channel + # bias before the activation, matching the device epilogue + # (D = act(alpha*acc + bias)). bias_storage holds the exact d_dtype values + # the kernel loads, so float() reproduces them bit-for-bit. + ref = alpha * ref + if use_bias: + ref = ref + bias_storage.float().reshape(1, 1, 1, 1, gemm_n).cuda() + # Add the per-element residual in FP32 before the activation, matching the + # device epilogue (D = act(alpha*acc + bias + beta*residual)). + # residual_storage holds the exact d_dtype values the kernel loads, so + # float() reproduces them bit-for-bit. + if has_residual: + ref = ref + beta * residual_storage.float().reshape(n, z, p, q, k).cuda() + # Mirror the device epilogue activation folded into the kernel. + ref = EPILOGUE_ACTIVATIONS[activation]["ref"](ref) + + d_ref_device = torch.empty((n, z, p, q, k), dtype=torch.float32, device="cuda") + cute.testing.convert( + output_, + from_dlpack(d_ref_device, assumed_align=16).mark_layout_dynamic( + leading_dim=4 + ), + ) + d_result = d_ref_device.cpu() + torch.testing.assert_close(d_result, ref.cpu(), atol=tolerance, rtol=1e-2) + print("Reference check passed.") + + def generate_tensors(): + # Only A/B/D rotate per cold-L2 workspace; the SF tensors, alpha, bias, + # and residual are reused from the outer scope (small / not perf-relevant + # to rotate), matching the golden benchmark. + input_tensor, filter_tensor, output_tensor = prepare_tensors( + ncdhw, ktrs, (z, p, q) + ) + input_, input_ws = create_cute_tensor(input_tensor, ab_dtype, leading_dim=4) + filter_, filter_ws = create_cute_tensor(filter_tensor, ab_dtype, leading_dim=4) + output_, output_ws = create_cute_tensor(output_tensor, d_dtype, leading_dim=4) + jit_args = testing.JitArguments( + input_, + filter_, + sfa_tensor, + sfb_tensor, + output_, + bias_, + residual_, + alpha_, + *rt_conv_params, + current_stream, + ) + # A from_dlpack tensor keeps its backing torch storage alive through the + # DLPack capsule it holds, so this pin is redundant; it stays as an + # explicit statement of the workspace's intended lifetime. + references = [input_ws, filter_ws, output_ws] + jit_args.add_to_scope(references) + # The workspace initializes on other streams; drain them so the + # benchmark stream never reads a tensor mid-initialization. + torch.cuda.synchronize() + return jit_args + + workspace_count = 1 + if use_cold_l2: + one_workspace_bytes = ( + input_storage.numel() * input_storage.element_size() + + filter_storage.numel() * filter_storage.element_size() + + output_storage.numel() * output_storage.element_size() + + sfa_storage.numel() * sfa_storage.element_size() + + sfb_storage.numel() * sfb_storage.element_size() + ) + # beta*residual is output-sized, so leaving it out shrinks the workspace + # ring below the real working set and lets L2 stay warm -- an optimistic + # cold-L2 measurement. bias is a small K vector but counted for symmetry. + if bias_storage is not None: + one_workspace_bytes += bias_storage.numel() * bias_storage.element_size() + if residual_storage is not None: + one_workspace_bytes += ( + residual_storage.numel() * residual_storage.element_size() + ) + workspace_count = testing.get_workspace_count( + one_workspace_bytes, warmup_iterations, iterations + ) + + if iterations > 0: + exec_time = testing.benchmark( + compiled_conv, + workspace_generator=generate_tensors, + workspace_count=workspace_count, + stream=current_stream, + warmup_iterations=warmup_iterations, + iterations=iterations, + use_cuda_graphs=True, + ) + gflop = 2 * gemm_m * gemm_n * gemm_k / 1e9 + gflops = gflop / exec_time * 1e6 + print(f"Execution time: {exec_time} microseconds per iteration") + print(f"GFLOPS: {gflops}") + return exec_time + return 0 + + +def run( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + stride_dhw: Tuple[int, int, int] = (1, 1, 1), + upper_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + lower_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + dil_dhw: Tuple[int, int, int] = (1, 1, 1), + ab_dtype: Type[cutlass.Numeric] = cutlass.Float4E2M1FN, + sf_dtype: Type[cutlass.Numeric] = cutlass.Float8E4M3FN, + sf_vec_size: int = 16, + d_dtype: Type[cutlass.Numeric] = cutlass.Float16, + acc_dtype: Type[cutlass.Numeric] = cutlass.Float32, + tile_shape_mnk: Tuple[int, int, int] = (128, 128, 128), + tolerance: float = 1e-1, + warmup_iterations: int = 0, + iterations: int = 1, + skip_ref_check: bool = False, + use_cold_l2: bool = False, + activation: str = "identity", + use_bias: bool = False, + alpha: float = 1.0, + beta: float = 0.0, + **kwargs, +): + return run_conv( + ncdhw, + ktrs, + ab_dtype, + sf_dtype, + sf_vec_size, + d_dtype, + acc_dtype, + tile_shape_mnk, + tolerance, + warmup_iterations, + iterations, + skip_ref_check, + use_cold_l2, + stride_dhw, + upper_pad_dhw, + lower_pad_dhw, + dil_dhw, + activation, + use_bias, + alpha, + beta, + ) + + +def parse_conv_arguments() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="SM120 NVFP4 static-specialized TRS fprop" + ) + parser.add_argument( + "--ncdhw", type=_parse_comma_separated_ints, default=(1, 128, 3, 3, 130) + ) + parser.add_argument( + "--ktrs", type=_parse_comma_separated_ints, default=(128, 3, 3, 3) + ) + parser.add_argument( + "--stride_dhw", type=_parse_comma_separated_ints, default=(1, 1, 1) + ) + parser.add_argument( + "--upper_pad_dhw", type=_parse_comma_separated_ints, default=(1, 1, 1) + ) + parser.add_argument( + "--lower_pad_dhw", type=_parse_comma_separated_ints, default=(1, 1, 1) + ) + parser.add_argument( + "--dil_dhw", type=_parse_comma_separated_ints, default=(1, 1, 1) + ) + parser.add_argument( + "--tile_shape_mnk", + type=_parse_comma_separated_ints, + choices=[ + (128, 128, 128), + (128, 96, 128), + (128, 64, 128), + (128, 128, 64), + (128, 96, 64), + (128, 64, 64), + ], + default=(128, 128, 128), + ) + parser.add_argument("--ab_dtype", type=cutlass.dtype, default=cutlass.Float4E2M1FN) + parser.add_argument("--sf_dtype", type=cutlass.dtype, default=cutlass.Float8E4M3FN) + parser.add_argument("--sf_vec_size", type=int, choices=[16, 32], default=16) + parser.add_argument("--d_dtype", type=cutlass.dtype, default=cutlass.Float16) + parser.add_argument("--acc_dtype", type=cutlass.dtype, default=cutlass.Float32) + parser.add_argument("--tolerance", type=float, default=1e-1) + parser.add_argument("--warmup_iterations", type=int, default=0) + parser.add_argument("--iterations", type=int, default=1) + parser.add_argument("--skip_ref_check", action="store_true", default=False) + parser.add_argument("--use_cold_l2", action="store_true", default=False) + parser.add_argument( + "--activation", + type=str, + default="identity", + choices=sorted(EPILOGUE_ACTIVATIONS), + help="Compile-time epilogue activation applied as D = activation(acc). " + "One activation per cubin.", + ) + parser.add_argument("--use_bias", action="store_true", default=False) + parser.add_argument( + "--alpha", + type=float, + default=1.0, + help="Per-tensor accumulator scale (D = act(alpha*acc + bias + " + "beta*residual)). Runtime scalar: one cubin serves any value.", + ) + parser.add_argument( + "--beta", + type=float, + default=0.0, + help="Residual scaling (D = act(alpha*acc + bias + beta*residual)); " + "beta == 0 disables residual, beta != 0 enables it. Compile-time " + "constant: one cubin per beta value.", + ) + args = parser.parse_args() + if len(args.ncdhw) != 5: + parser.error("--ncdhw must contain exactly 5 values") + if len(args.ktrs) != 4: + parser.error("--ktrs must contain exactly 4 values") + return args + + +if __name__ == "__main__": + args = parse_conv_arguments() + run( + args.ncdhw, + args.ktrs, + args.stride_dhw, + args.upper_pad_dhw, + args.lower_pad_dhw, + args.dil_dhw, + args.ab_dtype, + args.sf_dtype, + args.sf_vec_size, + args.d_dtype, + args.acc_dtype, + args.tile_shape_mnk, + args.tolerance, + args.warmup_iterations, + args.iterations, + args.skip_ref_check, + args.use_cold_l2, + activation=args.activation, + use_bias=args.use_bias, + alpha=args.alpha, + beta=args.beta, + ) + print("PASS") diff --git a/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/conv/dense_implicit_gemm_fprop.py b/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/conv/dense_implicit_gemm_fprop.py new file mode 100644 index 0000000000..6389ab16cb --- /dev/null +++ b/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/conv/dense_implicit_gemm_fprop.py @@ -0,0 +1,2584 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause + +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: + +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. + +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. + +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. + +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +import argparse +from typing import Optional, Tuple, Type +import sys +import os + +import cuda.bindings.driver as cuda + +import torch +import torch.nn.functional as F + +import cutlass +import cutlass.cute as cute +from cutlass import testing +from cutlass.cute.runtime import from_dlpack +import cutlass.torch as cutlass_torch +from cutlass.torch import dtype as torch_dtype +import cutlass.utils as utils +from cutlass.cute.nvgpu import cpasync +import cutlass.pipeline as pipeline +import cutlass.utils.hopper_helpers as sm90_utils +from cutlass.cute.arch.constants import WARPS_PER_WARPGROUP +from pathlib import Path + +if __name__ == "__main__": + # `helpers` sits at the examples/CuTeDSL root; running this file as a + # script only puts its own directory on sys.path. + cutedsl_dir = str(Path(__file__).resolve().parents[4]) + if cutedsl_dir not in sys.path: + sys.path.insert(0, cutedsl_dir) + +from helpers.dynamic_persistent_tile_scheduler import ( + ClcDynamicPersistentTileScheduler, + ClcDynamicPersistentTileSchedulerParams, +) + +if __name__ == "__main__": + current_dir = os.path.dirname(os.path.abspath(__file__)) + sys.path.insert(0, os.path.join(current_dir, "../../../")) + +""" +A high-performance 3D implicit-GEMM based fprop convolution example for the NVIDIA Blackwell Geforce +(SM120) architecture using CUTE DSL. +- Input tensor A is NxDxHxWxC, must be C major. +- Filter tensor B is KxTxRxSxC, must be C major. +- Output tensor D is NxZxPxQxK, must be K major. + +This kernel supports the following features: + - Utilizes Tensor Memory Access (TMA) im2col mode for efficient input loading with on-the-fly + im2col transformation + - Utilizes warp MMA for matrix multiply-accumulate (MMA) operations + - Supports multi-stage pipeline to overlap computation and memory access + - Ping-pong MMA warpgroups with parity-based tile ownership + - CLC-based dynamic persistent tile scheduling + - Optional per-output-channel bias and a compile-time epilogue activation, + D = activation(acc + bias) + +This implicit-GEMM based convolution works by converting the convolution into a GEMM problem: +- GEMM M dimension maps to NxZxPxQ +- GEMM N dimension maps to K +- GEMM K dimension maps to TxRxSxC +During the load of input tensor to SMEM, the TMA operation performs the im2col transformation +on the input tensor A. Filter tensor can be loaded to SMEM without any transformation. +The output tensor D is stored to GMEM via TMA. + +To run this example: + +.. code-block:: bash + + python dense_implicit_gemm_fprop.py \\ + --ncdhw 1,64,8,8,8 --ktrs 128,3,3,3 \\ + --tile_shape_mnk 128,128,64 \\ + --ab_dtype Float16 --d_dtype Float16 --acc_dtype Float32 \\ + --upper_pad_dhw 1,1,1 --lower_pad_dhw 1,1,1 \\ + --stride_dhw 1,1,1 --dil_dhw 1,1,1 + +Constraints: +* Supported input data types: fp16, bf16, fp8 (e4m3, e5m2) +* A/B tensor must have the same data type +* Supported accumulator data types: fp16, fp32 +* Supported output data types: fp16, bf16, fp32 +* CTA tile shape M must be 32 or a multiple of 64: a whole number of the 32-row MMA + tile, and a whole number of the epilogue tile, which caps at 64 rows +* CTA tile shape N must be a multiple of 32, the width of one MMA tile +* CTA tile shape K must be 32/64/128. Prefer the one whose trailing partial K tile + wastes fewest channels. +* The contiguous dimension of A/B/D tensors must be at least 16 bytes aligned, + i.e, number of elements is a multiple of 8 for Float16/BFloat16 and of 16 for + Float8E4M3FN/Float8E5M2. +* The convolution geometry must leave a positive output extent, so that the problem + covers at least one CTA tile along M and along N. +* The bias, when supplied, is a length-K tensor of the output data type. +""" + + +def _check_tensor_alignment( + c: int, + k: int, + ab_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], +): + """Check if the tensor alignment is valid for convolution.""" + + def check_contiguous_16B_alignment(dtype, num_major_elements): + num_contiguous_elements = 16 * 8 // dtype.width + return num_major_elements % num_contiguous_elements == 0 + + if not check_contiguous_16B_alignment( + d_dtype, k + ) or not check_contiguous_16B_alignment(ab_dtype, c): + raise testing.CantImplementError( + f"Invalid tensor alignment: C = {c}, K = {k}, ab_dtype = {ab_dtype}, d_dtype = {d_dtype}" + ) + + +def _check_im2col_descriptor_limits( + filter_trs: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + dil_dhw: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], +) -> None: + """Rejects a convolution geometry the im2col tensor map cannot encode. + + :param filter_trs: Filter extents (T, R, S) + :param stride_dhw: Convolution stride per spatial dimension + :param dil_dhw: Dilation per spatial dimension + :param upper_padding_dhw: Upper padding per spatial dimension + :param lower_padding_dhw: Lower padding per spatial dimension + + :raises testing.CantImplementError: If a field cannot hold its value + """ + # Fields of the 5D im2col tensor map are narrower than the convolution + # parameters feeding them, and a 3D convolution always builds a 5D + # descriptor, so those widths bind here. They come from the descriptor + # encoding: + # + # pixel-box corners one 5-bit signed field per spatial dimension, so + # [-16, 15] on W, H and D alike. The field holds both + # corners, and its width follows the rank: the 16 bits + # the encoding spends on corners are split across the + # rank - 2 spatial dimensions, giving one 16-bit field + # at rank 3, two 8-bit fields at rank 4, and three + # 5-bit fields here. + # traversal strides 3 bits holding the stride minus one, so [1, 8] + # + # A corner is not the padding itself -- the leading one is -lower_padding and + # the trailing one is upper_padding - (filter - 1) * dilation -- so padding + # and dilation trade against each other and only the combination is bounded. + # A corner past its field truncates into range and moves the pixel box, so the + # bound is enforced on the combination rather than on either input alone. + corner_lo = -16 + corner_hi = 15 + max_element_stride = 8 + for dim, flt, dil, pad_up, pad_lo, stride in zip( + ("D", "H", "W"), + filter_trs, + dil_dhw, + upper_padding_dhw, + lower_padding_dhw, + stride_dhw, + strict=True, + ): + leading_corner = -pad_lo + trailing_corner = pad_up - (flt - 1) * dil + if not corner_lo <= leading_corner <= corner_hi: + raise testing.CantImplementError( + f"{dim} leading im2col corner is -lower_padding = " + f"{leading_corner}, outside the [{corner_lo}, {corner_hi}] the " + f"descriptor's signed 5-bit corner encodes; lower_padding_" + f"{dim.lower()} must be at most {-corner_lo}" + ) + if not corner_lo <= trailing_corner <= corner_hi: + raise testing.CantImplementError( + f"{dim} trailing im2col corner is upper_padding - (filter - 1) * " + f"dilation = {pad_up} - ({flt} - 1) * {dil} = {trailing_corner}, " + f"outside the [{corner_lo}, {corner_hi}] the descriptor's signed " + f"5-bit corner encodes" + ) + if stride > max_element_stride: + raise testing.CantImplementError( + f"stride_{dim.lower()}={stride} exceeds the element stride of " + f"{max_element_stride} a tensor map allows" + ) + # A filter offset is dilation * tap index, so the largest one a dimension + # reaches is (filter - 1) * dilation. It travels in an unsigned 5-bit + # field per spatial dimension, so [0, 31] on W, H and D alike. Past it the + # value wraps into 5 bits rather than saturating -- an offset of 32 + # arrives as 0, collapsing every tap onto one position -- which changes + # the result without any fault, so it is bounded here even when both + # corners are in range. + # + # This bound and the trailing corner share the (filter - 1) * dilation + # term, so raising dilation moves both and a measurement that leaves + # upper_padding at zero cannot tell which one it broke: the corner gives + # out first, at an offset of 16. Raising upper_padding with the dilation + # holds the corner at -16 and isolates this bound. + max_filter_offset = 31 + filter_offset = (flt - 1) * dil + if filter_offset > max_filter_offset: + raise testing.CantImplementError( + f"{dim} filter offset is (filter - 1) * dilation = ({flt} - 1) * " + f"{dil} = {filter_offset}, past the {max_filter_offset} the " + f"im2col descriptor encodes for {dim}" + ) + + +def _compute_im2col_params( + filter_trs: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + dilation_dhw: Tuple[int, int, int], +) -> Tuple[ + Tuple[int, int, int], + Tuple[int, int, int], + Tuple[int, int, int], + Tuple[int, int, int], + Tuple[int, int, int], + Tuple[int, int, int], + Tuple[int, int, int], +]: + """Compute im2col TMA descriptor parameters from convolution parameters. + + Converts convolution parameters (DHW order) to TMA descriptor parameters (WHD order). + + :returns: (lower_corner_whd, upper_corner_whd, lower_padding_whd, + upper_padding_whd, stride_whd, lower_srt, stride_srt) + """ + pad_upper_d, pad_upper_h, pad_upper_w = upper_padding_dhw + pad_lower_d, pad_lower_h, pad_lower_w = lower_padding_dhw + stride_d, stride_h, stride_w = stride_dhw + dilation_d, dilation_h, dilation_w = dilation_dhw + filter_t, filter_r, filter_s = filter_trs + + lower_corner_whd = (-pad_lower_w, -pad_lower_h, -pad_lower_d) + upper_corner_whd = ( + pad_upper_w - ((filter_s - 1) * dilation_w), + pad_upper_h - ((filter_r - 1) * dilation_h), + pad_upper_d - ((filter_t - 1) * dilation_d), + ) + lower_padding_whd = (pad_lower_w, pad_lower_h, pad_lower_d) + upper_padding_whd = (pad_upper_w, pad_upper_h, pad_upper_d) + stride_whd = (stride_w, stride_h, stride_d) + lower_srt = (0, 0, 0) + stride_srt = (dilation_w, dilation_h, dilation_d) + + return ( + lower_corner_whd, + upper_corner_whd, + lower_padding_whd, + upper_padding_whd, + stride_whd, + lower_srt, + stride_srt, + ) + + +class Sm120PersistentDenseImplicitGemmFpropKernel: + """ + 3D convolution kernel for SM120 (Blackwell GeForce) via implicit GEMM. + The input (A) is expected in 5D tensor (NDHWC) format. + The filter (B) is expected in 5D tensor (KTRSC) format. + The output (D) is expected in 5D tensor (NZPQK) format. + + Ping-pong MMA warpgroups + CLC-based dynamic persistent scheduling. + + :param acc_dtype: Data type for accumulation during computation + :type acc_dtype: type[cutlass.Numeric] + :param tile_shape_mnk: CTA tile shape (M, N, K) + :type tile_shape_mnk: Tuple[int, int, int] + + Convolution geometry (filter T/R/S, padding, stride, dilation) is not a + constructor parameter: T/R/S come from the filter tensor extents and + pad/stride/dilation arrive as runtime Int32 operands to __call__, so a + single compiled cubin serves any geometry. + """ + + def __init__( + self, + acc_dtype: Type[cutlass.Numeric], + tile_shape_mnk: Tuple[int, int, int], + ): + self.acc_dtype = acc_dtype + self.cluster_shape_mnk = (1, 1, 1) + self.tile_shape_mnk = tuple(tile_shape_mnk) + self.tiled_mma = None + self.num_mcast_ctas_a = None + self.num_mcast_ctas_b = None + self.is_a_mcast = False + self.is_b_mcast = False + + self.occupancy = 1 + self.atom_layout = (2, 2, 1) + + # Use 2 mma warpgroups for ping pong + self.num_mma_warps = ( + self.atom_layout[0] * self.atom_layout[1] * self.atom_layout[2] * 2 + ) + self.num_dma_warps = 1 + self.num_sched_warps = 1 # CLC scheduler warp for dynamic persistent scheduling + self.num_threads_per_warp = 32 + self.threads_per_cta = ( + self.num_mma_warps + self.num_dma_warps + self.num_sched_warps + ) * self.num_threads_per_warp + + # Round up to the nearest warp group so that register reallocation works + self.threads_per_cta = (self.threads_per_cta + 127) // 128 * 128 + + self.smem_capacity = cutlass.memory.get_smem_capacity_in_bytes("sm_120") + + self.ab_stage = None + self.epi_stage = None + self.num_clc_stage = 1 + + self.a_smem_layout_staged = None + self.b_smem_layout_staged = None + self.epi_smem_layout_staged = None + self.epi_tile = None + + self.shared_storage = None + self.buffer_align_bytes = 1024 + + self.epilog_sync_barrier = pipeline.NamedBarrier( + barrier_id=2, + num_threads=128, + ) + # Bias staging ring depth: two rows let the DMA warp stage the next + # tile's bias row while the current tile's owner still reads its own. + self.bias_stage = 2 + self.load_register_requirement = 40 + self.mma_register_requirement = 232 + + def _setup_conv_input_attrs(self, a, b, d): + """Validate and set input-dependent attributes for convolution. + + :param a: Input tensor A - (N, D, H, W, C) layout + :param b: Filter tensor B - (K, T, R, S, C) layout + :param d: Output tensor D - (N, Z, P, Q, K) layout + """ + self.a_dtype = a.element_type + self.b_dtype = b.element_type + self.d_dtype = d.element_type + + # Only C-contiguous accepted for A/B, K-contiguous for C + if cutlass.const_expr(a.leading_dim != 4): + raise RuntimeError( + "The layout of a is not supported (must be C-contiguous)" + ) + if cutlass.const_expr(b.leading_dim != 4): + raise RuntimeError( + "The layout of b is not supported (must be C-contiguous)" + ) + if cutlass.const_expr(d.leading_dim != 4): + raise RuntimeError( + "The layout of d is not supported (must be K-contiguous)" + ) + + # K-major layouts for all operands in the GEMM view + self.a_layout = cutlass.tensor_utils.LayoutEnum.ROW_MAJOR + self.b_layout = cutlass.tensor_utils.LayoutEnum.ROW_MAJOR + self.d_layout = cutlass.tensor_utils.LayoutEnum.ROW_MAJOR + + # Check if input data types are compatible + if cutlass.const_expr(self.a_dtype != self.b_dtype): + raise TypeError(f"Type must match: {self.a_dtype} != {self.b_dtype}") + + def _setup_conv_tma(self, a, b, d, upper_pad_op, lower_pad_op, stride_op, dil_op): + """Set up TMA atoms and tensors for im2col convolution. + + The pad/stride/dilation operands feed the im2col A descriptor corners. + Threading them as runtime cutlass.Int32 lets one compiled cubin serve + any pad/stride/dilation config. + + :param a: Input tensor A - (N, D, H, W, C) layout + :param b: Filter tensor B - (K, T, R, S, C) layout + :param d: Output tensor D - (N, Z, P, Q, K) layout + :param upper_pad_op: Upper padding (D, H, W) as runtime cutlass.Int32 tuple + :param lower_pad_op: Lower padding (D, H, W) as runtime cutlass.Int32 tuple + :param stride_op: Convolution stride (D, H, W) as runtime cutlass.Int32 tuple + :param dil_op: Dilation (D, H, W) as runtime cutlass.Int32 tuple + :returns: (tma_atom_a, tma_tensor_a, tma_atom_b, tma_tensor_b, + tma_atom_d, tma_tensor_d) + """ + # Filter T/R/S sourced from the filter tensor b (K, T, R, S, C): under + # the dynamic tensor layout these extents are runtime Int32, so one + # compiled cubin serves any T/R/S. + rt_filter_trs = (b.shape[1], b.shape[2], b.shape[3]) + + # Compute im2col descriptor parameters (DHW -> WHD order) + ( + lower_corner_whd, + upper_corner_whd, + lower_padding_whd, + upper_padding_whd, + stride_whd, + lower_srt, + stride_srt, + ) = _compute_im2col_params( + rt_filter_trs, + upper_pad_op, + lower_pad_op, + stride_op, + dil_op, + ) + + # --- A: im2col TMA load --- + a_copy_atom = cpasync.CopyBulkTensorIm2ColG2SOp() + + # Create 2-mode hierarchical tensor layout: (N, D, H, W, C) -> ((W, H, D, N), C) + mA = cute.make_tensor(a.iterator, cute.select(a.layout, mode=[3, 2, 1, 0, 4])) + mA = cute.group_modes(mA, begin=0, end=4) + + a_smem_layout = cute.slice_(self.a_smem_layout_staged, (None, None, 0)) + + # Use the lower-level make_im2col_tma_atom (compatible with SM90-style rank-2 SMEM layouts) + tma_atom_a, tma_tensor_a = cpasync.make_im2col_tma_atom( + a_copy_atom, + mA, + a_smem_layout, + # cta_tiler = (M, (K,)): the MMA tiler tiles the channel dimension, so + # the K mode is nested to match the grouped channel mode of mA. + (self.tile_shape_mnk[0], (self.tile_shape_mnk[2],)), + lower_corner_whd, + upper_corner_whd, + lower_padding_whd, + upper_padding_whd, + stride_whd, + lower_srt, + stride_srt, + ) + + # --- B: tiled TMA load (filter reshaped to 2D) --- + # Change view of filter tensor from (K, T, R, S, C) to (K, (C, S, R, T)) + mB = cute.make_tensor(b.iterator, cute.select(b.layout, mode=[0, 4, 3, 2, 1])) + mB = cute.group_modes(mB, begin=1, end=5) + + tma_atom_b, tma_tensor_b = self._make_tma_atoms_and_tensors( + mB, + self.b_smem_layout_staged, + (self.tile_shape_mnk[1], (self.tile_shape_mnk[2],)), + 1, + ) + + # --- C: TMA im2col store --- + # Change view of output tensor from (N, Z, P, Q, K) to ((Q, P, Z, N), K) + mD = cute.make_tensor(d.iterator, cute.select(d.layout, mode=[3, 2, 1, 0, 4])) + mD = cute.group_modes(mD, begin=0, end=4) + + epi_smem_layout = cute.slice_(self.epi_smem_layout_staged, (None, None, 0)) + + tma_atom_d, tma_tensor_d = cpasync.make_im2col_tma_atom( + cpasync.CopyBulkTensorIm2ColS2GOp(), + mD, + epi_smem_layout, + self.epi_tile, + ) + tma_tensor_d = cute.coalesce(tma_tensor_d, target_profile=(1, 1)) + + # Add dummy batch dimension to all tensors (GEMM kernel expects batch dimension) + def add_dummy_batch_dimension(tensor): + new_layout = cute.append(tensor.layout, cute.make_layout(1)) + tensor = cute.make_tensor(tensor.iterator, new_layout) + return tensor + + tma_tensor_a = add_dummy_batch_dimension(tma_tensor_a) + tma_tensor_b = add_dummy_batch_dimension(tma_tensor_b) + tma_tensor_d = add_dummy_batch_dimension(tma_tensor_d) + + return ( + tma_atom_a, + tma_tensor_a, + tma_atom_b, + tma_tensor_b, + tma_atom_d, + tma_tensor_d, + ) + + def _setup_attributes(self): + # FP8 operands take the .e4m3/.e5m2 MMA on its K=32 atom, which is where + # FP8's doubled rate lives: one instruction covers twice the K. + is_fp8 = self.a_dtype in (cutlass.Float8E4M3FN, cutlass.Float8E5M2) + if is_fp8: + self.mma_inst_mnk = (16, 8, 32) + op = cute.nvgpu.warp.MmaFP8Op( + self.a_dtype, + self.acc_dtype, + self.mma_inst_mnk, + ) + else: + self.mma_inst_mnk = (16, 8, 16) + op = cute.nvgpu.warp.MmaF16BF16Op( + self.a_dtype, + self.acc_dtype, + self.mma_inst_mnk, + ) + + tC = cute.make_layout(self.atom_layout) + permutation_mnk = ( + self.atom_layout[0] * self.mma_inst_mnk[0], + self.atom_layout[1] * self.mma_inst_mnk[1] * 2, + self.atom_layout[2] * self.mma_inst_mnk[2], + ) + self.tiled_mma = cute.make_tiled_mma( + op, + tC, + permutation_mnk=permutation_mnk, + ) + + self.cta_layout_mnk = cute.make_layout(self.cluster_shape_mnk) + + self.num_mcast_ctas_a = self.cluster_shape_mnk[1] + self.num_mcast_ctas_b = self.cluster_shape_mnk[0] + self.is_a_mcast = self.num_mcast_ctas_a > 1 + self.is_b_mcast = self.num_mcast_ctas_b > 1 + + self.epi_tile = sm90_utils.compute_tile_shape_or_override( + self.tile_shape_mnk, self.d_dtype, is_cooperative=False + ) + + # Compute stage before compute smem layout + self.ab_stage, self.epi_stage = self._compute_stages( + self.tile_shape_mnk, + self.a_dtype, + self.b_dtype, + self.epi_tile, + self.d_dtype, + self.smem_capacity, + self.occupancy, + ) + + if self.ab_stage < 1: + raise testing.CantImplementError( + f"CTA tile {self.tile_shape_mnk} leaves shared memory for " + f"{self.ab_stage} A/B stages; it needs at least one" + ) + + ( + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.epi_smem_layout_staged, + ) = self._make_smem_layouts( + self.tile_shape_mnk, + self.epi_tile, + self.a_dtype, + self.a_layout, + self.b_dtype, + self.b_layout, + self.ab_stage, + self.d_dtype, + self.d_layout, + self.epi_stage, + ) + + @cute.jit + def __call__( + self, + a: cute.Tensor, + b: cute.Tensor, + d: cute.Tensor, + bias: Optional[cute.Tensor], + rt_upper_pad_d: cutlass.Int32, + rt_upper_pad_h: cutlass.Int32, + rt_upper_pad_w: cutlass.Int32, + rt_lower_pad_d: cutlass.Int32, + rt_lower_pad_h: cutlass.Int32, + rt_lower_pad_w: cutlass.Int32, + rt_stride_d: cutlass.Int32, + rt_stride_h: cutlass.Int32, + rt_stride_w: cutlass.Int32, + rt_dil_d: cutlass.Int32, + rt_dil_h: cutlass.Int32, + rt_dil_w: cutlass.Int32, + max_active_clusters: cutlass.Constexpr, + stream: cuda.CUstream, + epilogue_op: cutlass.Constexpr = lambda x: x, + ): + """Execute the convolution operation. + + :param a: Input tensor A - (N, D, H, W, C) layout + :param b: Filter tensor B - (K, T, R, S, C) layout + :param d: Output tensor D - (N, Z, P, Q, K) layout + :param bias: Optional per-output-channel bias of length K, or None. Whether + one is supplied is a compile-time branch, so a bias run and a no-bias run + are separate cubins + :param rt_upper_pad_d/h/w: Runtime upper padding (D, H, W) as Int32, so one + compiled cubin serves any padding config without recompilation + :param rt_lower_pad_d/h/w: Runtime lower padding (D, H, W) as Int32 + :param rt_stride_d/h/w: Runtime convolution stride (D, H, W) as Int32 + :param rt_dil_d/h/w: Runtime dilation (D, H, W) as Int32 + :param max_active_clusters: Maximum active clusters for scheduling + :param stream: CUDA stream for asynchronous execution + :param epilogue_op: Activation applied to the FP32 accumulator, folded into + the cubin as a Constexpr, so each activation compiles to its own kernel + """ + # Validate and set input-dependent attributes + self._setup_conv_input_attrs(a, b, d) + + # Bias opt-in: the caller supplies a length-K (output-channel) tensor. has_bias + # is a compile-time branch, so a no-bias run compiles the bias path away. + self.has_bias: bool = bias is not None + # The bias register fragment is allocated with the bias tensor's dtype and + # upconverted to FP32 for the add, so the bias has to match the output type. + if cutlass.const_expr(self.has_bias and bias.element_type is not self.d_dtype): + raise testing.CantImplementError( + f"bias dtype ({bias.element_type}) must match output dtype " + f"({self.d_dtype})" + ) + + # Setup attributes (tiled_mma, smem layouts, stages, etc.) + self._setup_attributes() + + # Pack the runtime pad/stride/dilation scalars into (D, H, W) tuples that + # feed the im2col A descriptor corners. Keeping them as runtime Int32 lets + # a single compiled cubin run any pad/stride/dilation configuration. + upper_pad_op = (rt_upper_pad_d, rt_upper_pad_h, rt_upper_pad_w) + lower_pad_op = (rt_lower_pad_d, rt_lower_pad_h, rt_lower_pad_w) + stride_op = (rt_stride_d, rt_stride_h, rt_stride_w) + dil_op = (rt_dil_d, rt_dil_h, rt_dil_w) + + # Create im2col TMA atoms and tensors + ( + tma_atom_a, + tma_tensor_a, + tma_atom_b, + tma_tensor_b, + tma_atom_d, + tma_tensor_d, + ) = self._setup_conv_tma(a, b, d, upper_pad_op, lower_pad_op, stride_op, dil_op) + + # Build the mBias tensor: a per-output-channel bias broadcast to the same + # (M, N, L) profile as the output. The bias varies along the output channel K + # (GEMM-N, real stride 1) and broadcasts across every spatial output position + # (GEMM-M carries stride 0), so the epilogue reads it through the same + # partition_C chain as the accumulator and each thread lands on its own + # N-column bias scalar. A static M extent keeps the partitioned smem-read + # layout static while the N axis carries a runtime extent, so one cubin serves + # any output-channel count. + if cutlass.const_expr(self.has_bias): + mBias_layout = cute.make_layout( + (self.tile_shape_mnk[0], cute.size(d, mode=[4]), 1), + stride=(0, 1, 0), + ) + mBias_mnl = cute.make_tensor(bias.iterator, mBias_layout) + else: + mBias_mnl = None + + # Compute grid from the reshaped C tensor (with dummy batch dim) + tile_sched_params, grid = self._compute_grid( + tma_tensor_d, + self.tile_shape_mnk, + max_active_clusters, + ) + + @cute.struct + class SharedStorage: + mainloop_pipeline_array_ptr: cute.struct.MemRange[ + cutlass.Int64, self.ab_stage * 2 + ] + # 2 Stage * 2 Groups order pipeline + pingpong_pipeline_array_ptr: cute.struct.MemRange[cutlass.Int64, 4] + clc_mbar_ptr: cute.struct.MemRange[cutlass.Int64, self.num_clc_stage * 2] + clc_response: cute.struct.Align[ + cute.struct.MemRange[cutlass.Int32, self.num_clc_stage * 4], 16 + ] + sA: cute.struct.Align[ + cute.struct.MemRange[ + self.a_dtype, cute.cosize(self.a_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + sB: cute.struct.Align[ + cute.struct.MemRange[ + self.b_dtype, cute.cosize(self.b_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + sD: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, cute.cosize(self.epi_smem_layout_staged) + ], + self.buffer_align_bytes, + ] + # Ring of cta_tile_n bias rows: the DMA warp cp.async's each tile's + # row into the stage its producer state points at, and the tile's + # owning math warpgroup reads it back through its consumer state. + # Zero length without a bias so it costs no smem. + sBias: cute.struct.Align[ + cute.struct.MemRange[ + self.d_dtype, + self.bias_stage * self.tile_shape_mnk[1] if self.has_bias else 0, + ], + self.buffer_align_bytes, + ] + # PipelineCpAsync mbarriers for the bias staging ring (full+empty per + # stage). Zero when no bias is supplied. + bias_pipeline_array_ptr: cute.struct.MemRange[ + cutlass.Int64, self.bias_stage * 2 if self.has_bias else 0 + ] + + self.shared_storage = SharedStorage + + self.kernel( + tma_atom_a, + tma_tensor_a, + tma_atom_b, + tma_tensor_b, + tma_atom_d, + tma_tensor_d, + mBias_mnl, + self.tiled_mma, + self.cta_layout_mnk, + self.a_smem_layout_staged, + self.b_smem_layout_staged, + self.epi_smem_layout_staged, + tile_sched_params, + epilogue_op, + ).launch( + grid=grid, + block=[self.threads_per_cta, 1, 1], + cluster=[1, 1, 1], + stream=stream, + min_blocks_per_mp=1, + ) + return + + # GPU device kernel + @cute.kernel + def kernel( + self, + tma_atom_a: cute.CopyAtom, + mA_mkl: cute.Tensor, + tma_atom_b: cute.CopyAtom, + mB_nkl: cute.Tensor, + tma_atom_d: cute.CopyAtom, + mD_mnl: cute.Tensor, + mBias_mnl: Optional[cute.Tensor], + tiled_mma: cute.TiledMma, + cta_layout_mnk: cute.Layout, + a_smem_layout_staged: cute.ComposedLayout, + b_smem_layout_staged: cute.ComposedLayout, + epi_smem_layout_staged: cute.ComposedLayout, + tile_sched_params: ClcDynamicPersistentTileSchedulerParams, + epilogue_op: cutlass.Constexpr, + ): + """ + GPU device kernel performing the batched GEMM computation. + + :param tma_atom_a: TMA copy atom for A tensor + :type tma_atom_a: cute.CopyAtom + :param mA_mkl: Input tensor A + :type mA_mkl: cute.Tensor + :param tma_atom_b: TMA copy atom for B tensor + :type tma_atom_b: cute.CopyAtom + :param mB_nkl: Input tensor B + :type mB_nkl: cute.Tensor + :param tma_atom_d: TMA copy atom for C tensor + :type tma_atom_d: cute.CopyAtom + :param mD_mnl: Output tensor D + :type mD_mnl: cute.Tensor + :param tiled_mma: Tiled MMA object + :type tiled_mma: cute.TiledMma + :param cta_layout_mnk: CTA layout + :type cta_layout_mnk: cute.Layout + :param a_smem_layout_staged: Shared memory layout for A + :type a_smem_layout_staged: cute.ComposedLayout + :param b_smem_layout_staged: Shared memory layout for B + :type b_smem_layout_staged: cute.ComposedLayout + :param epi_smem_layout_staged: Shared memory layout for epilogue + :type epi_smem_layout_staged: cute.ComposedLayout + """ + + # /////////////////////////////////////////////////////////////////////////////// + # Get cta/warp/thread idx + # /////////////////////////////////////////////////////////////////////////////// + tidx, _, _ = cute.arch.thread_idx() + warp_idx = cute.arch.warp_idx() + warp_idx = cute.arch.make_warp_uniform(warp_idx) + bidx, bidy, bidz = cute.arch.block_idx() + + # ///////////////////////////////////////////////////////////////////////////// + # Prefetch Tma desc + # ///////////////////////////////////////////////////////////////////////////// + if warp_idx == 0: + cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_a) + cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_b) + cute.nvgpu.cpasync.prefetch_descriptor(tma_atom_d) + + cta_rank_in_cluster = cute.arch.make_warp_uniform( + cute.arch.block_idx_in_cluster() + ) + cluster_coord_mnk = cta_layout_mnk.get_flat_coord(cta_rank_in_cluster) + + # /////////////////////////////////////////////////////////////////////////////// + # Get mcast mask + # /////////////////////////////////////////////////////////////////////////////// + a_mcast_mask = cute.make_layout_image_mask( + cta_layout_mnk, cluster_coord_mnk, mode=1 + ) + b_mcast_mask = cute.make_layout_image_mask( + cta_layout_mnk, cluster_coord_mnk, mode=0 + ) + + a_mcast_mask = a_mcast_mask if self.is_a_mcast else 0 + b_mcast_mask = b_mcast_mask if self.is_b_mcast else 0 + a_smem_layout = cute.slice_(a_smem_layout_staged, (None, None, 0)) + b_smem_layout = cute.slice_(b_smem_layout_staged, (None, None, 0)) + tma_copy_bytes = cute.size_in_bytes( + self.a_dtype, a_smem_layout + ) + cute.size_in_bytes(self.b_dtype, b_smem_layout) + + # ///////////////////////////////////////////////////////////////////////////// + # Alloc and init AB full/empty + ACC full mbar (pipeline) + # ///////////////////////////////////////////////////////////////////////////// + smem = cutlass.memory.SmemAllocator() + storage = smem.allocate(self.shared_storage) + + # mbar arrays + mainloop_pipeline_array_ptr = storage.mainloop_pipeline_array_ptr.data_ptr() + pingpong_pipeline_array_ptr = storage.pingpong_pipeline_array_ptr.data_ptr() + clc_mbar_ptr = storage.clc_mbar_ptr.data_ptr() + clc_response_ptr = storage.clc_response.data_ptr() + + # Threads/warps participating in this pipeline + mainloop_pipeline_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread + ) + mainloop_pipeline_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Warp, WARPS_PER_WARPGROUP + ) + + cta_layout_vmnk = cute.make_layout((1, *cta_layout_mnk.shape)) + mainloop_pipeline = pipeline.PipelineTmaAsync.create( + num_stages=self.ab_stage, + producer_group=mainloop_pipeline_producer_group, + consumer_group=mainloop_pipeline_consumer_group, + tx_count=tma_copy_bytes, + barrier_storage=mainloop_pipeline_array_ptr, + cta_layout_vmnk=cta_layout_vmnk, + enable_multicast_signaling=True, + defer_sync=True, + ) + + pingpong_pipeline_group = pipeline.CooperativeGroup(pipeline.Agent.Thread, 128) + warp_group_idx = cute.arch.make_warp_uniform(tidx // 128) + pingpong_pipeline = pipeline.PipelineOrder.create( + barrier_storage=pingpong_pipeline_array_ptr, + depth=2, + length=2, + group_id=warp_group_idx, + producer_group=pingpong_pipeline_group, + defer_sync=True, + ) + + # CLC pipeline: + # Consumers are the sched warp (self-consume for exit detection), the DMA warp, + # and both MMA warpgroups. Each MMA warpgroup picks up every tile and uses a + # parity counter to skip the ones owned by the other warpgroup. + clc_producer_group = pipeline.CooperativeGroup(pipeline.Agent.Thread) + clc_consumer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + (self.num_mma_warps + self.num_dma_warps + self.num_sched_warps) + * self.num_threads_per_warp, + ) + clc_pipeline = pipeline.PipelineClcFetchAsync.create( + num_stages=self.num_clc_stage, + producer_group=clc_producer_group, + consumer_group=clc_consumer_group, + tx_count=16, + barrier_storage=clc_mbar_ptr, + cta_layout_vmnk=cta_layout_vmnk, + defer_sync=True, + ) + + # Bias staging pipeline: the DMA warp cp.async's each tile's bias row + # into the smem ring and the tile's owning math warpgroup reads it back. + # The producer commit is a per-lane cp.async arrive from the DMA warp's + # 32 threads; the release is a per-thread arrive from the owning + # warpgroup's 128 threads, so each stage's participant set is fixed no + # matter which warpgroup owns the tile. + if cutlass.const_expr(mBias_mnl is not None): + bias_pipeline = pipeline.PipelineCpAsync.create( + num_stages=self.bias_stage, + producer_group=pipeline.CooperativeGroup( + pipeline.Agent.Thread, self.num_threads_per_warp + ), + consumer_group=pipeline.CooperativeGroup(pipeline.Agent.Thread, 128), + barrier_storage=storage.bias_pipeline_array_ptr.data_ptr(), + defer_sync=True, + ) + else: + bias_pipeline = None + + cute.arch.mbarrier_init_fence() + # Cluster arrive after barrier init + if cute.size(self.cluster_shape_mnk) > 1: + cute.arch.cluster_arrive_relaxed() + + # /////////////////////////////////////////////////////////////////////////////// + # Generate smem tensor A/B + # /////////////////////////////////////////////////////////////////////////////// + sA = storage.sA.get_tensor( + a_smem_layout_staged.outer, swizzle=a_smem_layout_staged.inner + ) + sB = storage.sB.get_tensor( + b_smem_layout_staged.outer, swizzle=b_smem_layout_staged.inner + ) + # (cta_tile_n, STAGE) ring of bias rows: the DMA warp fills the stage + # its producer state points at, the tile's owning math warpgroup reads + # the stage its consumer state points at. + if cutlass.const_expr(mBias_mnl is not None): + cta_n = self.tile_shape_mnk[1] + sBias = storage.sBias.get_tensor( + cute.make_layout((cta_n, self.bias_stage), stride=(1, cta_n)) + ) + else: + cta_n = None + sBias = None + + sD = storage.sD.get_tensor( + epi_smem_layout_staged.outer, swizzle=epi_smem_layout_staged.inner + ) + + # /////////////////////////////////////////////////////////////////////////////// + # Local_tile partition global tensors + # /////////////////////////////////////////////////////////////////////////////// + # (bM, bK, loopM, loopK, loopL) + gA_mkl = cute.local_tile( + mA_mkl, + (self.tile_shape_mnk[0], (self.tile_shape_mnk[2],)), + (None, None, None), + ) + # (bN, bK, loopN, loopK, loopL) + gB_nkl = cute.local_tile( + mB_nkl, + (self.tile_shape_mnk[1], (self.tile_shape_mnk[2],)), + (None, None, None), + ) + # (bM, bN, loopM, loopN, loopL) + gD_mnl = cute.local_tile( + mD_mnl, + cute.slice_(self.tile_shape_mnk, (None, None, 0)), + (None, None, None), + ) + # Bias shares the output's MNL tiling; its M axis carries a stride-0 + # broadcast so every spatial output row reads the same per-channel bias. + if cutlass.const_expr(mBias_mnl is not None): + gBias_mnl = cute.local_tile( + mBias_mnl, + cute.slice_(self.tile_shape_mnk, (None, None, 0)), + (None, None, None), + ) + else: + gBias_mnl = None + + # ////////////////////////////////////////////////////////////////////////////// + # Partition global tensor for TiledMMA_A/B/C + # ////////////////////////////////////////////////////////////////////////////// + thr_mma = tiled_mma.get_slice(tidx % 128) + + # ////////////////////////////////////////////////////////////////////////////// + # Partition shared tensor for TMA load A/B + # ////////////////////////////////////////////////////////////////////////////// + # TMA load A partition_S/D + a_cta_layout = cute.make_layout(cute.slice_(cta_layout_mnk, (0, None, 0)).shape) + a_cta_crd = cluster_coord_mnk[1] + tAsA, tAgA = cute.nvgpu.cpasync.tma_partition( + tma_atom_a, + a_cta_crd, + a_cta_layout, + cute.group_modes(sA, 0, 2), + cute.group_modes(gA_mkl, 0, 2), + ) + + # TMA load B partition_S/D + b_cta_layout = cute.make_layout(cute.slice_(cta_layout_mnk, (None, 0, 0)).shape) + b_cta_crd = cluster_coord_mnk[0] + tBsB, tBgB = cute.nvgpu.cpasync.tma_partition( + tma_atom_b, + b_cta_crd, + b_cta_layout, + cute.group_modes(sB, 0, 2), + cute.group_modes(gB_nkl, 0, 2), + ) + + # Make frangments + tCsA = thr_mma.partition_A(sA) + tCsB = thr_mma.partition_B(sB) + tCrA = tiled_mma.make_fragment_A(tCsA[None, None, None, 0]) + tCrB = tiled_mma.make_fragment_B(tCsB[None, None, None, 0]) + + tDgD = thr_mma.partition_C(gD_mnl) + acc_shape = tDgD.shape[:3] + accumulators = cute.make_rmem_tensor(acc_shape, self.acc_dtype) + # Bias runs the same partition_C chain as the accumulator, so the fragment + # read back from smem lines up element for element with the acc fragment. + if cutlass.const_expr(mBias_mnl is not None): + # Identity coordinates through the mma C partition, so each fragment + # element carries its own (m, n). The n coordinate is what indexes this + # warpgroup's sBias row and what gates the N overhang. + cBias_mnl = cute.make_identity_tensor(gBias_mnl.shape) + tCcBias = thr_mma.partition_C(cBias_mnl) + # cp.async transfers at least 32 bits, so each lane moves a 32-bit vector + # of bias elements: two at the output's 16-bit width. + bias_elems_per_copy = 32 // mBias_mnl.element_type.width + bias_g2s_atom = cute.make_copy_atom( + cute.nvgpu.cpasync.CopyG2SOp(), + mBias_mnl.element_type, + num_bits_per_copy=32, + ) + else: + tCcBias = None + bias_elems_per_copy = None + bias_g2s_atom = None + + # cluster wait for barrier init + if cute.size(self.cluster_shape_mnk) > 1: + cute.arch.cluster_wait() + else: + pipeline.sync(barrier_id=1) + + k_tile_cnt = cute.size(gA_mkl, mode=[3]) + + # Create the tile scheduler (CLC-based dynamic persistent scheduler). + tile_sched = ClcDynamicPersistentTileScheduler.create( + tile_sched_params, + cute.arch.block_idx(), + cute.arch.grid_dim(), + clc_response_ptr, + ) + work_tile = tile_sched.initial_work_tile_info() + + # Create the pipeline states for producer and consumer + mainloop_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.ab_stage + ) + mainloop_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.ab_stage + ) + + clc_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.num_clc_stage + ) + + # MMA warp groups + if warp_idx < self.num_mma_warps: + cute.arch.setmaxregister_increase(self.mma_register_requirement) + + num_k_blocks = cute.size(tCrA, mode=[2]) + + # /////////////////////////////////////////////////////////////////////////////// + # Copy Atom A/B retiling for TMA load A/B + # /////////////////////////////////////////////////////////////////////////////// + atom_copy_ldmatrix_A = cute.make_copy_atom( + cute.nvgpu.warp.LdMatrix8x8x16bOp(self.a_layout.is_m_major_a(), 4), + self.a_dtype, + ) + atom_copy_ldmatrix_B = cute.make_copy_atom( + cute.nvgpu.warp.LdMatrix8x8x16bOp(self.b_layout.is_n_major_b(), 4), + self.b_dtype, + ) + smem_tiled_copy_A = cute.make_tiled_copy_A(atom_copy_ldmatrix_A, tiled_mma) + + smem_tiled_copy_B = cute.make_tiled_copy_B(atom_copy_ldmatrix_B, tiled_mma) + + thr_copy_ldmatrix_A = smem_tiled_copy_A.get_slice(tidx % 128) + thr_copy_ldmatrix_B = smem_tiled_copy_B.get_slice(tidx % 128) + tCsA_copy_view = thr_copy_ldmatrix_A.partition_S(sA) + tCrA_copy_view = thr_copy_ldmatrix_A.retile(tCrA) + + tCsB_copy_view = thr_copy_ldmatrix_B.partition_S(sB) + tCrB_copy_view = thr_copy_ldmatrix_B.retile(tCrB) + + # Parity counter: wg0 owns tiles at parity 0 (tile 0, tile 2, tile 4, ...), + # wg1 owns parity 1. The warpgroup whose parity does NOT match skips the + # tile and advances its mainloop_consumer_state by k_tile_cnt so DMA's linear production stays + # in phase with this warpgroup's next owned tile. + tile_parity = cutlass.Int32(0) + + # Bias ring consumer state. The DMA warp produces one stage per tile, + # so the non-owner advances past a skipped tile's stage to keep the + # ring counter aligned with its next owned tile. + if cutlass.const_expr(mBias_mnl is not None): + bias_consumer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Consumer, self.bias_stage + ) + + pingpong_pipeline_state = pingpong_pipeline.state + while work_tile.is_valid_tile: + if tile_parity == warp_group_idx: + pingpong_pipeline.wait(pingpong_pipeline_state) + tile_coord_mnl = work_tile.tile_idx + + gD_mnl_slice = gD_mnl[(None, None, *tile_coord_mnl)] + # Clear the accumulator + accumulators.fill(0.0) + + # ///////////////////////////////////////////////////////////////////////////// + # Pipelined MAINLOOP + # ///////////////////////////////////////////////////////////////////////////// + + mainloop_consumer_state.reset_count() + + peek_ab_full_status = cutlass.Boolean(1) + if mainloop_consumer_state.count < k_tile_cnt: + peek_ab_full_status = mainloop_pipeline.consumer_try_wait( + mainloop_consumer_state + ) + + # Wait for TMA copies to complete + mainloop_pipeline.consumer_wait( + mainloop_consumer_state, peek_ab_full_status + ) + # tCsA_p: (MMA, (4, MMA_M / 4), MMA_K), tCsA_p: (MMA, (4, MMA_N / 4), MMA_K) + tCsA_p = tCsA_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + tCsB_p = tCsB_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + cute.copy( + smem_tiled_copy_A, + tCsA_p[None, None, 0], + tCrA_copy_view[None, None, 0], + ) + cute.copy( + smem_tiled_copy_B, + tCsB_p[None, None, 0], + tCrB_copy_view[None, None, 0], + ) + + for k_tile in range(0, k_tile_cnt - 1, 1, unroll=1): + # unroll the loop + for k_block_idx in cutlass.range_constexpr(num_k_blocks): + k_block_next = ( + 0 + if k_block_idx + 1 == num_k_blocks + else k_block_idx + 1 + ) + + if k_block_idx == num_k_blocks - 1: + cute.arch.fence_view_async_shared() + mainloop_pipeline.consumer_release( + mainloop_consumer_state + ) + mainloop_consumer_state.advance() + + peek_ab_full_status = cutlass.Boolean(1) + peek_ab_full_status = ( + mainloop_pipeline.consumer_try_wait( + mainloop_consumer_state + ) + ) + + # tCsA_p: (MMA, (4, MMA_M / 4), MMA_K), tCsA_p: (MMA, (4, MMA_N / 4), MMA_K) + tCsA_p = tCsA_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + tCsB_p = tCsB_copy_view[ + None, None, None, mainloop_consumer_state.index + ] + mainloop_pipeline.consumer_wait( + mainloop_consumer_state, peek_ab_full_status + ) + + # Copy data from smem to tCrA/tCrB for the next k_block. + # Above one k_block the prefetch targets a different register + # slot than the MMA is about to read, so issuing it first + # overlaps the load. At exactly one, k_block_next wraps onto + # that same slot after the tile was released, so there it has + # to come after the MMA. + def prefetch_next_k_block(): + cute.copy( + smem_tiled_copy_A, + tCsA_p[None, None, k_block_next], + tCrA_copy_view[None, None, k_block_next], + ) + cute.copy( + smem_tiled_copy_B, + tCsB_p[None, None, k_block_next], + tCrB_copy_view[None, None, k_block_next], + ) + + if cutlass.const_expr(num_k_blocks > 1): + prefetch_next_k_block() + # Gemm of the current k_block + cute.gemm( + tiled_mma, + accumulators, + tCrA[None, None, k_block_idx], + tCrB[None, None, k_block_idx], + accumulators, + ) + if cutlass.const_expr(num_k_blocks == 1): + prefetch_next_k_block() + # end of for loop + # Hoist out last k_tile + + for k_block_idx in cutlass.range_constexpr(num_k_blocks): + k_block_next = ( + 0 if k_block_idx + 1 == num_k_blocks else k_block_idx + 1 + ) + + if k_block_idx == num_k_blocks - 1: + cute.arch.fence_view_async_shared() + + mainloop_pipeline.consumer_release(mainloop_consumer_state) + mainloop_consumer_state.advance() + + # Signal other warpgroup to proceed + pingpong_pipeline_state = pingpong_pipeline.arrive( + pingpong_pipeline_state + ) + + if k_block_next > 0: + cute.copy( + smem_tiled_copy_A, + tCsA_p[None, None, k_block_next], + tCrA_copy_view[None, None, k_block_next], + ) + cute.copy( + smem_tiled_copy_B, + tCsB_p[None, None, k_block_next], + tCrB_copy_view[None, None, k_block_next], + ) + # Gemm of the current k_block + cute.gemm( + tiled_mma, + accumulators, + tCrA[None, None, k_block_idx], + tCrB[None, None, k_block_idx], + accumulators, + ) + + # Add the per-output-channel bias in FP32 (D = act(acc + bias)). + # The DMA warp has cp.async'd this tile's contiguous cta_tile_n + # bias values into the ring stage this warpgroup's consumer + # state points at; every thread loads its own N columns back + # and adds them to the accumulator. One value per output + # channel is fetched once and broadcast to all M rows by the + # stride-0 M axis of mBias_mnl. + if cutlass.const_expr(mBias_mnl is not None): + bias_pipeline.consumer_wait(bias_consumer_state) + sBias_row = sBias[(None, bias_consumer_state.index)] + # Read back into a fragment aligned with the accumulator. + # The identity-coordinate partition gives each element its + # tile-local N column in [0, cta_n), which is exactly the + # linear index into the staged row: the producer already + # applied n_base when staging from gmem and zero-filled any + # overhang. Indexing the row by that tile-local column, + # rather than by a gmem channel stride, is what makes the + # read correct. + tCcBias_tile = tCcBias[(None, None, None, *tile_coord_mnl)] + tCrBias = cute.make_rmem_tensor( + accumulators.shape, mBias_mnl.element_type + ) + for be in cutlass.range_constexpr(cute.size(tCrBias)): + tCrBias[be] = sBias_row[tCcBias_tile[be][1]] + # Release the stage so the DMA warp can refill it: every + # thread arrives only after its smem reads above are done. + cute.arch.fence_proxy("async.shared", space="cta") + bias_pipeline.consumer_release(bias_consumer_state) + bias_consumer_state.advance() + accumulators.store( + accumulators.load() + tCrBias.load().to(self.acc_dtype) + ) + + # ///////////////////////////////////////////////////////////////////////////// + # EPILOG + # ///////////////////////////////////////////////////////////////////////////// + + copy_atom_r2s = sm90_utils.sm90_get_smem_store_op( + self.d_layout, + elem_ty_d=self.d_dtype, + elem_ty_acc=self.acc_dtype, + ) + + # StMatrix is a 16-bit instruction; pass f16 so the partition + # geometry is consistent regardless of d_dtype. + # The actual rmem->smem store is performed by copy_atom_r2s above. + copy_atom_C = cute.make_copy_atom( + cute.nvgpu.warp.StMatrix8x8x16bOp( + self.d_layout.is_m_major_c(), + 4, + ), + cutlass.Float16, + ) + + tiled_copy_C_Atom = cute.make_tiled_copy_C_atom( + copy_atom_C, tiled_mma + ) + + tiled_copy_r2s = cute.make_tiled_copy_S( + copy_atom_r2s, + tiled_copy_C_Atom, + ) + + thr_copy_r2s = tiled_copy_r2s.get_slice(tidx % 128) + # (R2S, R2S_M, R2S_N, PIPE_D) + tRS_sD = thr_copy_r2s.partition_D(sD) + # (R2S, R2S_M, R2S_N) + tRS_rAcc = tiled_copy_r2s.retile(accumulators) + + # Allocate D registers. + rD_shape = cute.shape(thr_copy_r2s.partition_S(sD)) + tRS_rD_layout = cute.make_layout(rD_shape[:3]) + tRS_rD = cute.make_rmem_tensor(tRS_rD_layout.shape, self.acc_dtype) + + # One epi tile covers a sub-range of the accumulator's M and N + # modes. The accumulator linearises as v + V * (m + M * n), so a + # tile that spans several n values skips the whole m extent between + # them and its elements are not contiguous. How many n values a + # tile spans follows the output width -- a narrower element leaves + # room for a wider epi tile -- so the block has to be addressed by + # its (m, n) coordinate rather than by a flat offset. + acc_m_cnt = cute.size(tRS_rAcc, mode=[1]) + rD_v_cnt = cute.size(tRS_rD, mode=[0]) + rD_m_cnt = cute.size(tRS_rD, mode=[1]) + rD_n_cnt = cute.size(tRS_rD, mode=[2]) + + sepi_for_tma_partition = cute.group_modes(sD, 0, 2) + tcgc_for_tma_partition = cute.zipped_divide( + gD_mnl_slice, self.epi_tile + ) + + bSG_sD, bSG_gD = cute.nvgpu.cpasync.tma_partition( + tma_atom_d, + 0, + cute.make_layout(1), + sepi_for_tma_partition, + tcgc_for_tma_partition, + ) + + epi_tile_num = cute.size(tcgc_for_tma_partition, mode=[1]) + epi_tile_shape = tcgc_for_tma_partition.shape[1] + epi_tile_layout = cute.make_layout( + epi_tile_shape, stride=(1, epi_tile_shape[0]) + ) + # Initialize tma store pipeline + tma_store_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.num_mma_warps * self.num_threads_per_warp, + ) + tma_store_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.epi_stage, + producer_group=tma_store_producer_group, + ) + + # Serialize epilogue between warp groups (guards shared sD) + pingpong_pipeline.wait(pingpong_pipeline_state) + + for epi_idx in cutlass.range_constexpr(epi_tile_num): + # Copy from accumulators to D registers + epi_m_blk, epi_n_blk = epi_tile_layout.get_hier_coord(epi_idx) + for epi_n in cutlass.range_constexpr(rD_n_cnt): + acc_n = epi_n_blk * rD_n_cnt + epi_n + for epi_m in cutlass.range_constexpr(rD_m_cnt): + acc_m = epi_m_blk * rD_m_cnt + epi_m + for epi_v in cutlass.range_constexpr(rD_v_cnt): + tRS_rD[ + epi_v + rD_v_cnt * (epi_m + rD_m_cnt * epi_n) + ] = tRS_rAcc[ + epi_v + rD_v_cnt * (acc_m + acc_m_cnt * acc_n) + ] + + # Apply the activation in FP32, then cast to the output + # type. The op is folded into the cubin as a Constexpr, so each + # activation compiles to its own kernel. + tRS_rD_out = cute.make_rmem_tensor( + tRS_rD_layout.shape, self.d_dtype + ) + acc_vec = epilogue_op(tRS_rD.load()) + tRS_rD_out.store(acc_vec.to(self.d_dtype)) + + # Register to shared memory + epi_buffer = epi_idx % cute.size(tRS_sD, mode=[3]) + cute.copy( + tiled_copy_r2s, + tRS_rD_out, + tRS_sD[(None, None, None, epi_buffer)], + ) + + cute.arch.fence_view_async_shared() + # barrier for sync + self.epilog_sync_barrier.arrive_and_wait() + + # Get the global memory coordinate for the current epi tile. + gmem_coord = epi_tile_layout.get_hier_coord(epi_idx) + # Copy from shared memory to global memory + if warp_idx % 4 == 0: + cute.copy( + tma_atom_d, + bSG_sD[(None, epi_buffer)], + bSG_gD[(None, gmem_coord)], + ) + tma_store_pipeline.producer_commit() + tma_store_pipeline.producer_acquire() + # barrier for sync + self.epilog_sync_barrier.arrive_and_wait() + + tma_store_pipeline.producer_tail() + # Signal other warpgroup it can start its epilogue + pingpong_pipeline_state = pingpong_pipeline.arrive( + pingpong_pipeline_state + ) + else: + # Manually advance pipeline stage by the k_tile_cnt stages DMA produced. + for k_tile in range(0, k_tile_cnt, 1, unroll=1): + mainloop_consumer_state.advance() + # Likewise skip this tile's bias stage: the DMA warp produced + # one per tile, so advance past it to stay aligned with the + # next owned tile. + if cutlass.const_expr(mBias_mnl is not None): + bias_consumer_state.advance() + + # Pull the next tile from the CLC response slot. Both MMA warpgroups, + # plus the DMA warp and the sched warp, participate as consumers. + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + tile_parity = cutlass.Int32(1) - tile_parity + # End of while work_tile.is_valid_tile + # End of MMA warp group + # Start of DMA warp group + elif warp_idx == self.num_mma_warps: + cute.arch.setmaxregister_decrease(self.load_register_requirement) + + if cutlass.const_expr(mBias_mnl is not None): + bias_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.bias_stage + ) + + while work_tile.is_valid_tile: + tile_coord_mnl = work_tile.tile_idx + + tAgA_mkl = tAgA[(None, tile_coord_mnl[0], None, tile_coord_mnl[2])] + tBgB_nkl = tBgB[(None, tile_coord_mnl[1], None, tile_coord_mnl[2])] + + mainloop_producer_state.reset_count() + k_shape = cute.shape(tAgA_mkl, mode=1) + coord_iter = cute.repeat_like(0, k_shape) + + for k_tile in range(0, k_tile_cnt, 1, unroll=1): + # ///////////////////////////////////////////////////////////////////////////// + # Wait for A/B buffers to be empty before loading into them + # Also sets the transaction barrier for the A/B buffers + # ///////////////////////////////////////////////////////////////////////////// + mainloop_pipeline.producer_acquire(mainloop_producer_state) + + # ///////////////////////////////////////////////////////////////////////////// + # Slice to global/shared memref to current k_tile + # ///////////////////////////////////////////////////////////////////////////// + tAgA_k = tAgA_mkl[(None, coord_iter)] + tAsA_pipe = tAsA[(None, mainloop_producer_state.index)] + + tBgB_k = tBgB_nkl[(None, coord_iter)] + tBsB_pipe = tBsB[(None, mainloop_producer_state.index)] + + # ///////////////////////////////////////////////////////////////////////////// + # TMA load A/B + # ///////////////////////////////////////////////////////////////////////////// + cute.copy( + tma_atom_a, + tAgA_k, + tAsA_pipe, + tma_bar_ptr=mainloop_pipeline.producer_get_barrier( + mainloop_producer_state + ), + mcast_mask=a_mcast_mask, + ) + cute.copy( + tma_atom_b, + tBgB_k, + tBsB_pipe, + tma_bar_ptr=mainloop_pipeline.producer_get_barrier( + mainloop_producer_state + ), + mcast_mask=b_mcast_mask, + ) + # Mainloop pipeline's producer commit is a NOP + mainloop_pipeline.producer_commit(mainloop_producer_state) + mainloop_producer_state.advance() + + coord_iter = cute.increment_coord(coord_iter, k_shape) + + # Stage this tile's bias row for its owning math warpgroup: wait + # for the ring stage to drain, cp.async the row in, and commit so + # each lane's arrive lands once its copies complete. + if cutlass.const_expr(mBias_mnl is not None): + bias_pipeline.producer_acquire(bias_producer_state) + bias_lane_idx = tidx % self.num_threads_per_warp + n_base = tile_coord_mnl[1] * cta_n + # Each lane cp.async's contiguous 32-bit vectors + # (bias_elems_per_copy elements each) of this tile's bias row. + # Column-major stride so lane t owns [t*elems, (t+1)*elems); + # the output-channel axis has stride 1, so each segment is + # contiguous. n_active lanes cover cta_n. + n_active = cta_n // bias_elems_per_copy + bias_row_layout = cute.make_layout( + (bias_elems_per_copy, n_active), + stride=(1, bias_elems_per_copy), + ) + # cp.async needs 32-bit source and destination alignment; the + # tile base is a multiple of cta_tile_n, which keeps both on a + # 4-byte boundary, so re-annotate the pointers to satisfy the + # 32-bit atom. + gBias_row = cute.make_tensor( + (mBias_mnl.iterator + n_base).align(min_align=4), + bias_row_layout, + ) + sBias_stage = sBias[(None, bias_producer_state.index)] + sBias_tiled = cute.make_tensor( + sBias_stage.iterator.align(min_align=4), bias_row_layout + ) + bias_pred = cute.make_rmem_tensor( + cute.make_layout((1,)), cutlass.Boolean + ) + # The warp has 32 lanes and the row needs n_active of them; + # each extra pass hands the lanes the next 32 vectors. + for bias_pass in cutlass.range_constexpr((n_active + 31) // 32): + bias_lane = bias_pass * 32 + bias_lane_idx + if bias_lane < n_active: + # A CTA N-tile rounds up to cta_tile_n, but the + # output-channel count need not divide it, so tail + # lanes address bias columns past the end with no + # backing storage. Guard each lane's vector on its + # base channel: in-bounds lanes copy from gmem, + # out-of-bounds lanes zero-fill (cp.async writes 0 on + # a false predicate). The zero tail is only read back + # for overhang output the TMA store clamps away. + bias_pred[0] = cutlass.Boolean( + n_base + bias_lane * bias_elems_per_copy + < mBias_mnl.shape[1] + ) + cute.copy_atom_call( + bias_g2s_atom, + gBias_row[(None, bias_lane)], + sBias_tiled[(None, bias_lane)], + pred=bias_pred, + ) + cute.arch.cp_async_commit_group() + bias_pipeline.producer_commit(bias_producer_state) + bias_producer_state.advance() + + # Pull the next tile from the CLC queue. + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + # end of while loop + + # Wait A/B buffer empty + mainloop_pipeline.producer_tail(mainloop_producer_state) + + # Start of CLC scheduler warp + elif warp_idx == self.num_mma_warps + self.num_dma_warps: + cute.arch.setmaxregister_decrease(self.load_register_requirement) + + clc_producer_state = pipeline.make_pipeline_state( + pipeline.PipelineUserType.Producer, self.num_clc_stage + ) + + while work_tile.is_valid_tile: + clc_pipeline.producer_acquire(clc_producer_state) + tile_sched.advance_to_next_work( + clc_pipeline.producer_get_barrier(clc_producer_state) + ) + clc_producer_state.advance() + + # Self-consume to learn whether the next tile is valid. + clc_pipeline.consumer_wait(clc_consumer_state) + work_tile = tile_sched.get_current_work() + clc_pipeline.consumer_release(clc_consumer_state) + clc_consumer_state.advance() + + clc_pipeline.producer_tail(clc_producer_state) + + # Unused warps + else: + cute.arch.setmaxregister_decrease(self.load_register_requirement) + + return + + @staticmethod + def _compute_stages( + tile_shape_mnk: tuple[int, int, int], + a_dtype: type[cutlass.Numeric], + b_dtype: type[cutlass.Numeric], + epi_tile: tuple[int, int], + d_dtype: type[cutlass.Numeric], + smem_capacity: int, + occupancy: int, + ) -> tuple[int, int]: + """Computes the number of stages for the A/B operands and the epilogue. + + :param tile_shape_mnk: The shape (M, N, K) of the CTA tile. + :type tile_shape_mnk: tuple[int, int, int] + :param a_dtype: Data type of operand A. + :type a_dtype: type[cutlass.Numeric] + :param b_dtype: Data type of operand B. + :type b_dtype: type[cutlass.Numeric] + :param epi_tile: Epilogue tile shape. + :type epi_tile: tuple[int, int] + :param d_dtype: Data type of the output D. + :type d_dtype: type[cutlass.Numeric] + :param smem_capacity: Total available shared memory capacity in bytes. + :type smem_capacity: int + :param occupancy: Target number of CTAs per SM (occupancy). + :type occupancy: int + + :return: A tuple containing the computed number of stages for: + (A/B operand stages, epilogue stages) + :rtype: tuple[int, int] + """ + epi_stage = 4 + d_bytes_per_stage = cute.size(epi_tile) * d_dtype.width // 8 + epi_bytes = d_bytes_per_stage * epi_stage + + a_shape = cute.slice_(tile_shape_mnk, (None, 0, None)) + b_shape = cute.slice_(tile_shape_mnk, (0, None, None)) + ab_bytes_per_stage = ( + cute.size(a_shape) * a_dtype.width // 8 + + cute.size(b_shape) * b_dtype.width // 8 + ) + mbar_helpers_bytes = 1024 + + ab_stage = ( + (smem_capacity - occupancy * 1024) // occupancy + - mbar_helpers_bytes + - epi_bytes + ) // ab_bytes_per_stage + return ab_stage, epi_stage + + @staticmethod + def _make_smem_layouts( + tile_shape_mnk: tuple[int, int, int], + epi_tile: tuple[int, int], + a_dtype: type[cutlass.Numeric], + a_layout: cute.Layout, + b_dtype: type[cutlass.Numeric], + b_layout: cute.Layout, + ab_stage: int, + d_dtype: type[cutlass.Numeric], + d_layout: cute.Layout, + epi_stage: int, + ) -> tuple[cute.ComposedLayout, cute.ComposedLayout, cute.ComposedLayout]: + """Create shared memory layouts for the A, B and D tensors. + + :param tile_shape_mnk: CTA tile shape (M, N, K) + :type tile_shape_mnk: tuple[int, int, int] + :param epi_tile: Epilogue tile shape + :type epi_tile: tuple[int, int] + :param a_dtype: Data type for matrix A + :type a_dtype: type[cutlass.Numeric] + :param a_layout: Layout for matrix A + :type a_layout: cute.Layout + :param b_dtype: Data type for matrix B + :type b_dtype: type[cutlass.Numeric] + :param b_layout: Layout for matrix B + :type b_layout: cute.Layout + :param ab_stage: Number of stages for the A/B tensors + :type ab_stage: int + :param d_dtype: Data type for the output matrix D + :type d_dtype: type[cutlass.Numeric] + :param d_layout: Layout for the output matrix D + :type d_layout: cute.Layout + :param epi_stage: Number of epilogue stages + :type epi_stage: int + + :return: Tuple of shared memory layouts for A, B and the epilogue + :rtype: Tuple[cute.ComposedLayout, cute.ComposedLayout, cute.ComposedLayout] + """ + a_smem_layout_staged = sm90_utils.make_smem_layout_a( + a_layout, + tile_shape_mnk, + a_dtype, + ab_stage, + ) + + b_smem_layout_staged = sm90_utils.make_smem_layout_b( + b_layout, + tile_shape_mnk, + b_dtype, + ab_stage, + ) + + epi_smem_layout_staged = sm90_utils.make_smem_layout_epi( + d_dtype, + d_layout, + epi_tile, + epi_stage, + ) + + return a_smem_layout_staged, b_smem_layout_staged, epi_smem_layout_staged + + @staticmethod + def _compute_grid( + d: cute.Tensor, + tile_shape_mnk: tuple[int, int, int], + max_active_clusters: cutlass.Constexpr, + ) -> tuple[int, int, int]: + """Compute grid shape for the output tensor D. + + :param d: The output tensor D + :type d: cute.Tensor + :param tile_shape_mnk: The shape (M, N, K) of the CTA tile. + :type tile_shape_mnk: tuple[int, int, int] + + :return: Grid shape for kernel launch. + :rtype: tuple[int, int, int] + """ + + d_shape = cute.slice_(tile_shape_mnk, (None, None, 0)) + gd = cute.zipped_divide(d, tiler=d_shape) + num_ctas_mnl = gd[(0, (None, None, None))].shape + cluster_shape_mnl = (1, 1, 1) + tile_sched_params = ClcDynamicPersistentTileSchedulerParams( + num_ctas_mnl, cluster_shape_mnl + ) + grid = ClcDynamicPersistentTileScheduler.get_grid_shape(tile_sched_params) + return tile_sched_params, grid + + def can_implement( + self, + gemm_m: int, + c: int, + k: int, + ab_dtype: Type[cutlass.Numeric], + acc_dtype: Type[cutlass.Numeric], + d_dtype: Type[cutlass.Numeric], + filter_trs: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + dil_dhw: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], + bias_dtype: Optional[Type[cutlass.Numeric]] = None, + ) -> None: + """Rejects configurations this kernel cannot compile or would miscompute. + + Every constraint here is decidable from the tile shape, the element types + and the problem extents, so it is checked before any tensor is allocated + rather than surfacing as a rejected launch, an out-of-bounds store or a + silently wrong result. + + :param gemm_m: Implicit-GEMM M extent, N*Z*P*Q + :type gemm_m: int + :param c: Input channel count + :type c: int + :param k: Output channel count, which is the implicit-GEMM N extent + :type k: int + :param ab_dtype: A/B element type + :type ab_dtype: Type[cutlass.Numeric] + :param acc_dtype: Accumulator element type + :type acc_dtype: Type[cutlass.Numeric] + :param d_dtype: Output element type + :type d_dtype: Type[cutlass.Numeric] + :param filter_trs: Filter extents (T, R, S) + :type filter_trs: Tuple[int, int, int] + :param stride_dhw: Convolution stride per spatial dimension + :type stride_dhw: Tuple[int, int, int] + :param dil_dhw: Dilation per spatial dimension + :type dil_dhw: Tuple[int, int, int] + :param upper_padding_dhw: Upper padding per spatial dimension + :type upper_padding_dhw: Tuple[int, int, int] + :param lower_padding_dhw: Lower padding per spatial dimension + :type lower_padding_dhw: Tuple[int, int, int] + :param bias_dtype: Bias element type, or None when no bias is supplied + :type bias_dtype: Optional[Type[cutlass.Numeric]] + + :raises testing.CantImplementError: If the configuration is unsupported + """ + allowed_ab_dtype = ( + cutlass.Float16, + cutlass.BFloat16, + cutlass.Float8E4M3FN, + cutlass.Float8E5M2, + ) + if ab_dtype not in allowed_ab_dtype: + raise testing.CantImplementError( + f"ab_dtype must be one of {allowed_ab_dtype}, got {ab_dtype}" + ) + allowed_acc_dtype = (cutlass.Float32, cutlass.Float16) + if acc_dtype not in allowed_acc_dtype: + raise testing.CantImplementError( + f"acc_dtype must be one of {allowed_acc_dtype}, got {acc_dtype}" + ) + allowed_d_dtype = ( + cutlass.Float8E4M3FN, + cutlass.Float8E5M2, + cutlass.Float16, + cutlass.BFloat16, + cutlass.Float32, + ) + if d_dtype not in allowed_d_dtype: + raise testing.CantImplementError( + f"d_dtype must be one of {allowed_d_dtype}, got {d_dtype}" + ) + + # tile_shape_mnk keeps K hierarchical as (M, N, (K,)), so M and N are plain + # extents and this runs on the host, before any kernel IR exists. + tile_m, tile_n = self.tile_shape_mnk[0], self.tile_shape_mnk[1] + + # The warp MMA covers 16 rows and 8 columns per instruction on every input + # type this kernel takes -- only its K depth follows the type -- and the + # tiled MMA lays atom_layout warps over that with the N side permuted twice + # as wide. The CTA tile has to be a whole number of those, so that every + # element of it is owned by a warp. + mma_tile_m = self.atom_layout[0] * 16 + mma_tile_n = self.atom_layout[1] * 8 * 2 + if tile_m % mma_tile_m != 0 or tile_n % mma_tile_n != 0: + raise testing.CantImplementError( + f"CTA tile ({tile_m}, {tile_n}) must be a whole number of " + f"({mma_tile_m}, {mma_tile_n}) MMA tiles" + ) + + # The epilogue walks the CTA tile in epilogue tiles, and that tile caps its + # rows at 64 rather than deriving them from the CTA tile, so above 64 rows the + # CTA tile has to be a whole number of them: a remainder step would store into + # the rows the next tile owns. The epilogue's column count caps at the MMA tile + # width for every output type here, so the check above covers N. + epi_max_rows = 64 + if tile_m > epi_max_rows and tile_m % epi_max_rows != 0: + raise testing.CantImplementError( + f"CTA tile M {tile_m} must be a multiple of the epilogue tile's " + f"{epi_max_rows} rows" + ) + + # The K tile is one of three widths, each a whole number of MMA K instructions + # at either atom width. Pick the one whose trailing partial K tile wastes the + # fewest channels, ceil(C / tile_k) * tile_k - C, with a tie going to the wider + # tile, as far as the width that still leaves more than one A/B stage: the + # narrowest spans a single MMA instruction and suits only a C no wider width + # divides, and the widest leaves the 16-bit operands a single A/B stage. + allowed_tile_k = (32, 64, 128) + tile_k = self.tile_shape_mnk[2] + if tile_k not in allowed_tile_k: + # FP8 takes the K=32 MMA atom; every other input type here takes K=16. + mma_inst_k = ( + 32 if ab_dtype in (cutlass.Float8E4M3FN, cutlass.Float8E5M2) else 16 + ) + raise testing.CantImplementError( + f"tile_k must be one of {allowed_tile_k}, got {tile_k} (the MMA K " + f"instruction is {mma_inst_k} channels wide)" + ) + + # The bias fragment is allocated with the bias tensor's element type and + # upconverted to FP32 for the add, so the bias carries the output's type. + if bias_dtype is not None and bias_dtype is not d_dtype: + raise testing.CantImplementError( + f"bias dtype must match output dtype ({d_dtype}), got {bias_dtype}" + ) + + _check_tensor_alignment(c, k, ab_dtype, d_dtype) + + # A cluster is built out of CTA tiles, so it needs at least as many tiles + # along a dimension as it has CTAs there. A trailing partial tile still + # occupies a full CTA in the grid, so the tiles are counted with a round-up. + # At the single-CTA cluster this kernel launches the floor is one tile per + # dimension, which is what rejects a geometry whose output extent came out at + # zero or below: the grid then has no tile to schedule and the launch is + # refused with cudaErrorInvalidValue. + m_tiles = -(-gemm_m // tile_m) + n_tiles = -(-k // tile_n) + cluster_mn = self.cluster_shape_mnk[:2] + if m_tiles < cluster_mn[0] or n_tiles < cluster_mn[1]: + raise testing.CantImplementError( + f"CTA tile count ({m_tiles}, {n_tiles}) from implicit-GEMM " + f"({gemm_m}, {k}) with cta_tile=({tile_m}, {tile_n}) cannot form one " + f"cluster of {cluster_mn} CTAs" + ) + + _check_im2col_descriptor_limits( + filter_trs, stride_dhw, dil_dhw, upper_padding_dhw, lower_padding_dhw + ) + + @staticmethod + def _make_tma_atoms_and_tensors( + tensor: cute.Tensor, + smem_layout_staged: cute.ComposedLayout, + smem_tile: tuple[int, int], + mcast_dim: int, + internal_type: Optional[Type[cutlass.Numeric]] = None, + ) -> tuple[cute.CopyAtom, cute.Tensor]: + """Create TMA atoms and tensors for input tensors. + + :param tensor: Input tensor (A or B) + :type tensor: cute.Tensor + :param smem_layout_staged: Shared memory layout for the tensor + :type smem_layout_staged: cute.ComposedLayout + :param smem_tile: Shared memory tile shape + :type smem_tile: Tuple[int, int] + :param mcast_dim: Multicast dimension + :type mcast_dim: int + :param internal_type: Element type the descriptor addresses, for a + sub-byte operand whose storage type it cannot encode + :type internal_type: Optional[Type[cutlass.Numeric]] + + :return: TMA atom and tensor + :rtype: Tuple[cute.CopyAtom, cute.Tensor] + """ + op = ( + cute.nvgpu.cpasync.CopyBulkTensorTileG2SOp() + if mcast_dim == 1 + else cute.nvgpu.cpasync.CopyBulkTensorTileG2SMulticastOp() + ) + + if cutlass.const_expr(cute.rank(smem_layout_staged) == 4): + smem_layout = cute.slice_(smem_layout_staged, (None, None, None, 0)) + else: + smem_layout = cute.slice_(smem_layout_staged, (None, None, 0)) + tma_atom, tma_tensor = cute.nvgpu.cpasync.make_tiled_tma_atom( + op, + tensor, + smem_layout, + smem_tile, + num_multicast=mcast_dim, + internal_type=internal_type, + ) + + return tma_atom, tma_tensor + + +# ///////////////////////////////////////////////////////////////////////////// +# Helper functions +# ///////////////////////////////////////////////////////////////////////////// + + +def compute_zpq( + dhw: Tuple[int, int, int], + trs: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int], + upper_padding_dhw: Tuple[int, int, int], + lower_padding_dhw: Tuple[int, int, int], + dilation_dhw: Tuple[int, int, int], +) -> Tuple[int, int, int]: + """Compute output spatial dimensions Z, P, and Q with asymmetric padding.""" + D, H, W = dhw + T, R, S = trs + Sd, Sh, Sw = stride_dhw + UpperPadD, UpperPadH, UpperPadW = upper_padding_dhw + LowerPadD, LowerPadH, LowerPadW = lower_padding_dhw + DilD, DilH, DilW = dilation_dhw + Z = ((D + UpperPadD + LowerPadD - DilD * (T - 1) - 1) // Sd) + 1 + P = ((H + UpperPadH + LowerPadH - DilH * (R - 1) - 1) // Sh) + 1 + Q = ((W + UpperPadW + LowerPadW - DilW * (S - 1) - 1) // Sw) + 1 + return Z, P, Q + + +def create_cute_tensor( + source_f32_tensor: torch.Tensor, + dtype: Type[cutlass.Numeric], + leading_dim: int = None, +) -> Tuple[cute.Tensor, torch.Tensor]: + """Create a dynamic-layout cute tensor from a source f32 tensor. + + The tensor is always marked dynamic-layout: its non-leading extents lower to + runtime SSA, so nothing about this tensor pins the kernel to a problem size. + + :param source_f32_tensor: Source f32 tensor + :param dtype: Element type of the tensor to build + :param leading_dim: Leading dimension for dynamic layout + :return: Tuple of cute tensor and storage tensor + """ + # FP4 needs packed storage: 2 elements per byte. cute_tensor_like cannot + # build the half-byte layout, so allocate an int8 buffer viewed as the + # packed fp4 type and convert the f32 source into it directly. + if dtype == cutlass.Float4E2M1FN: + shape = tuple(source_f32_tensor.shape) + if shape[-1] % 2 != 0: + raise ValueError(f"FP4 packed storage requires even trailing dim: {shape}") + packed_shape = shape[:-1] + (shape[-1] // 2,) + storage_int8 = torch.empty(packed_shape, dtype=torch.int8, device="cuda") + storage_view = storage_int8.view(dtype=torch.float4_e2m1fn_x2) + cute_tensor = from_dlpack(storage_view, assumed_align=16) + cute_tensor = cute_tensor.mark_layout_dynamic(leading_dim=leading_dim) + if source_f32_tensor.numel() > 0: + f32_tensor = from_dlpack(source_f32_tensor, assumed_align=16) + f32_tensor = f32_tensor.mark_layout_dynamic(leading_dim=leading_dim) + cute.testing.convert(f32_tensor, cute_tensor) + return cute_tensor, storage_view + + cute_tensor, storage_tensor = cutlass_torch.cute_tensor_like( + source_f32_tensor, + dtype, + is_dynamic_layout=True, + assumed_align=16, + ) + return cute_tensor, storage_tensor + + +def prepare_tensors( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + zpq: Tuple[int, int, int], +) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Prepare f32 tensors for 3D convolution. + + :param ncdhw: Input tensor shape (N, C, D, H, W) + :param ktrs: Filter tensor shape components (K, T, R, S) + :param zpq: Output spatial dimensions (Z, P, Q) + :return: Tuple of input, filter, and output tensors + """ + N, C, D, H, W = ncdhw + K, T, R, S = ktrs + Z, P, Q = zpq + + input_tensor = torch.randint( + -1, + 2, + (N, D, H, W, C), + dtype=torch.float32, + device="cuda", + ) + filter_tensor = torch.randint( + -1, + 2, + (K, T, R, S, C), + dtype=torch.float32, + device="cuda", + ) + output_tensor = torch.empty((N, Z, P, Q, K), dtype=torch.float32, device="cuda") + + return input_tensor, filter_tensor, output_tensor + + +def _rt_conv_scalars( + *, + upper_pad: Tuple[int, int, int], + lower_pad: Tuple[int, int, int], + stride: Tuple[int, int, int], + dil: Tuple[int, int, int], +) -> list: + """Box the runtime conv geometry as cutlass.Int32 in the kernel's argument order. + + The order (upper pad, lower pad, stride, dilation, each D/H/W) must match the + kernel __call__ signature; keyword-only parameters keep the four groups from + being swapped at a call site. + """ + return [ + cutlass.Int32(v) for group in (upper_pad, lower_pad, stride, dil) for v in group + ] + + +# Compile-time epilogue activations, folded into the cubin as a Constexpr op (one +# activation per cubin). Each entry pairs the device-side op applied to the FP32 +# output fragment with the torch op used to build the reference. +EPILOGUE_ACTIVATIONS = { + "identity": { + "device": lambda x: x, + "ref": lambda x: x, + }, + "relu": { + "device": lambda x: cute.where(x > 0, x, cute.full_like(x, 0)), + "ref": torch.nn.functional.relu, + }, +} + + +# ///////////////////////////////////////////////////////////////////////////// +# Run function +# ///////////////////////////////////////////////////////////////////////////// + + +def compile_conv( + input_: cute.Tensor, + filter_: cute.Tensor, + output_: cute.Tensor, + bias_: Optional[cute.Tensor], + acc_dtype: Type[cutlass.Numeric], + tile_shape_mnk: Tuple[int, int, int], + stride_dhw: Tuple[int, int, int] = (1, 1, 1), + upper_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + lower_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + dil_dhw: Tuple[int, int, int] = (1, 1, 1), + epilogue_op: cutlass.Constexpr = lambda x: x, +): + """Build the kernel object, resolve the host launch config, and cute.compile it. + + Returns the compiled callable. Pad/stride/dilation are boxed as cutlass.Int32 so + the compiled entry scalars lower to runtime SSA and stay out of the mangled kernel + name; one cubin then serves any of those values, and the caller re-boxes them at + launch time. The stream a compilation is handed never runs anything, so it is a + fake one; the caller passes the real stream when it launches. + """ + from cutlass.cute.runtime import make_fake_stream + + conv = Sm120PersistentDenseImplicitGemmFpropKernel( + acc_dtype=acc_dtype, + tile_shape_mnk=tile_shape_mnk, + ) + hardware_info = utils.HardwareInfo() + max_active_clusters = hardware_info.get_max_active_clusters(1) + + return cute.compile( + conv, + input_, + filter_, + output_, + bias_, + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + max_active_clusters, + make_fake_stream(), + epilogue_op, + ) + + +def run( + ncdhw: Tuple[int, int, int, int, int], + ktrs: Tuple[int, int, int, int], + stride_dhw: Tuple[int, int, int] = (1, 1, 1), + upper_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + lower_pad_dhw: Tuple[int, int, int] = (0, 0, 0), + dil_dhw: Tuple[int, int, int] = (1, 1, 1), + ab_dtype: Type[cutlass.Numeric] = cutlass.Float16, + d_dtype: Type[cutlass.Numeric] = cutlass.Float16, + acc_dtype: Type[cutlass.Numeric] = cutlass.Float32, + tile_shape_mnk: Tuple[int, int, int] = (128, 128, 64), + tolerance: float = 1e-01, + warmup_iterations: int = 0, + iterations: int = 1, + use_cold_l2: bool = False, + skip_ref_check: bool = False, + use_bias: bool = False, + activation: str = "identity", + **kwargs, +): + """Run 3D convolution and compare against PyTorch reference. + + :param ncdhw: Input tensor shape (N, C, D, H, W) + :param ktrs: Filter tensor shape components (K, T, R, S) + :param stride_dhw: Stride (Sd, Sh, Sw) + :param upper_pad_dhw: Upper padding (PadD, PadH, PadW) + :param lower_pad_dhw: Lower padding (PadD, PadH, PadW) + :param dil_dhw: Dilation (DilD, DilH, DilW) + :param ab_dtype: Data type for A/B input tensors + :param d_dtype: Data type for output tensor D + :param acc_dtype: Accumulator data type + :param tile_shape_mnk: CTA tile shape (M, N, K) + :param tolerance: Tolerance for result comparison + :param warmup_iterations: Number of warmup iterations + :param iterations: Number of benchmark iterations + :param use_cold_l2: Whether to flush L2 cache between iterations + :param skip_ref_check: Whether to skip reference checking + :param use_bias: Add a per-output-channel bias, D = activation(acc + bias). + Whether a bias is supplied is a compile-time branch, so this selects a + different cubin + :param activation: Name of the compile-time epilogue activation, one of + EPILOGUE_ACTIVATIONS. One activation per cubin + """ + if not torch.cuda.is_available(): + raise RuntimeError("GPU is required to run this example!") + + N, C, D, H, W = ncdhw + K, T, R, S = ktrs + + ab_dtype = getattr(cutlass, ab_dtype) if isinstance(ab_dtype, str) else ab_dtype + d_dtype = getattr(cutlass, d_dtype) if isinstance(d_dtype, str) else d_dtype + acc_dtype = getattr(cutlass, acc_dtype) if isinstance(acc_dtype, str) else acc_dtype + + Z, P, Q = compute_zpq( + (D, H, W), + (T, R, S), + stride_dhw, + upper_pad_dhw, + lower_pad_dhw, + dil_dhw, + ) + + print("Running Blackwell GeForce (SM120) 3D Convolution test with:") + print(f" Input shape (N, C, D, H, W): {ncdhw}") + print(f" Filter shape (K, C, T, R, S): ({K}, {C}, {T}, {R}, {S})") + print(f" Output shape (N, K, Z, P, Q): ({N}, {K}, {Z}, {P}, {Q})") + print(f" Stride (Sd, Sh, Sw): {stride_dhw}") + print(f" Upper padding (PadD, PadH, PadW): {upper_pad_dhw}") + print(f" Lower padding (PadD, PadH, PadW): {lower_pad_dhw}") + print(f" Dilation (DilD, DilH, DilW): {dil_dhw}") + print(f" A/B data type: {ab_dtype}") + print(f" D data type: {d_dtype}") + print(f" Accumulator type: {acc_dtype}") + print(f" Tile shape (M, N, K): {tile_shape_mnk}\n") + + # Create convolution kernel. Convolution geometry (T/R/S, pad, stride, + # dilation) is not baked into the kernel object: T/R/S come from the filter + # tensor extents and pad/stride/dilation are passed as runtime Int32 to the + # compiled function below. + # Resolve the compile-time epilogue activation. The device op is folded into the + # cubin, so a different activation is a different kernel. + if activation not in EPILOGUE_ACTIVATIONS: + raise testing.CantImplementError( + f"Unsupported activation {activation!r}; " + f"choose from {sorted(EPILOGUE_ACTIVATIONS)}" + ) + epilogue_op = EPILOGUE_ACTIVATIONS[activation]["device"] + + conv = Sm120PersistentDenseImplicitGemmFpropKernel( + acc_dtype=acc_dtype, + tile_shape_mnk=tile_shape_mnk, + ) + conv.can_implement( + N * Z * P * Q, + C, + K, + ab_dtype, + acc_dtype, + d_dtype, + (T, R, S), + stride_dhw, + dil_dhw, + upper_pad_dhw, + lower_pad_dhw, + d_dtype if use_bias else None, + ) + + # Per-output-channel bias of length K, in the output's dtype. Added to the + # accumulator in FP32 as D = activation(acc + bias) and broadcast across every + # spatial output position. + if use_bias: + bias_storage = torch.randn(K, dtype=torch.float32, device="cuda").to( + torch_dtype(d_dtype) + ) + bias_ = from_dlpack(bias_storage, assumed_align=16) + else: + bias_storage, bias_ = None, None + + # Create input and filter tensors + input_tensor, filter_tensor, output_tensor = prepare_tensors(ncdhw, ktrs, (Z, P, Q)) + + # Prepare cute tensors + input_, input_storage = create_cute_tensor(input_tensor, ab_dtype, leading_dim=4) + filter_, filter_storage = create_cute_tensor(filter_tensor, ab_dtype, leading_dim=4) + output_, output_storage = create_cute_tensor(output_tensor, d_dtype, leading_dim=4) + + print("Compiling kernel with cute.compile ...") + compiled_fn = compile_conv( + input_, + filter_, + output_, + bias_, + acc_dtype, + tile_shape_mnk, + stride_dhw=stride_dhw, + upper_pad_dhw=upper_pad_dhw, + lower_pad_dhw=lower_pad_dhw, + dil_dhw=dil_dhw, + epilogue_op=epilogue_op, + ) + + # Get current CUDA stream + torch_stream = torch.cuda.Stream() + current_stream = cuda.CUstream(torch_stream.cuda_stream) + + # Run convolution. Pad/stride/dilation are passed as runtime Int32 so one + # cubin runs any pad/stride/dilation config. + print("Running Blackwell GeForce 3D convolution...") + # The inputs initialize on other streams; drain them so the kernel's + # stream never reads a tensor mid-initialization. + torch.cuda.synchronize() + compiled_fn( + input_, + filter_, + output_, + bias_, + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + current_stream, + ) + torch_stream.synchronize() + + with torch.backends.cudnn.flags(enabled=False): + if not skip_ref_check: + # Run PyTorch reference convolution + print("Running PyTorch reference 3D convolution...") + # Pytorch expects tensors to be in NCDHW format + input_ncdhw = input_tensor.permute(0, 4, 1, 2, 3) + if upper_pad_dhw != lower_pad_dhw: + # F.conv3d only supports symmetric padding, so manually pad the input + # F.pad takes padding in reverse dimension order + pad_arg = ( + lower_pad_dhw[2], + upper_pad_dhw[2], + lower_pad_dhw[1], + upper_pad_dhw[1], + lower_pad_dhw[0], + upper_pad_dhw[0], + ) + input_ncdhw = F.pad(input_ncdhw, pad_arg) + conv_padding = (0, 0, 0) + else: + conv_padding = upper_pad_dhw + ref = F.conv3d( + input_ncdhw, + filter_tensor.permute(0, 4, 1, 2, 3), + stride=stride_dhw, + padding=conv_padding, + dilation=dil_dhw, + ) + # Add the bias and apply the activation in FP32 before the output cast, + # matching the device epilogue (D = activation(acc + bias)). + # bias_storage holds the exact d_dtype values the kernel loads, so float() + # reproduces them bit for bit. + if use_bias: + ref = ref + bias_storage.float().reshape(1, K, 1, 1, 1) + ref = EPILOGUE_ACTIVATIONS[activation]["ref"](ref) + output_ref = ref.to(dtype=torch_dtype(d_dtype)).to(dtype=torch.float32) + # Compare results + print("Comparing results...") + + # Transform output from (N, Z, P, Q, K) -> (N, K, Z, P, Q) + output_f32 = output_storage.permute(0, 4, 1, 2, 3).to(torch.float32) + output_ref_f32 = output_ref.to(torch.float32) + + torch.testing.assert_close( + output_f32, + output_ref_f32, + atol=tolerance, + rtol=1e-03, + ) + print("Results match within tolerance!") + + # Benchmark if requested + if iterations > 0: + print( + f"\nBenchmarking with {warmup_iterations} warmup and {iterations} iterations..." + ) + + def generate_tensors(): + input_tensor, filter_tensor, output_tensor = prepare_tensors( + ncdhw, ktrs, (Z, P, Q) + ) + input_, _ = create_cute_tensor( + input_tensor, + ab_dtype, + leading_dim=4, + ) + filter_, _ = create_cute_tensor( + filter_tensor, + ab_dtype, + leading_dim=4, + ) + output_, _ = create_cute_tensor( + output_tensor, + d_dtype, + leading_dim=4, + ) + # The workspace initializes on other streams; drain them so the + # benchmark stream never reads a tensor mid-initialization. + torch.cuda.synchronize() + return testing.JitArguments( + input_, + filter_, + output_, + bias_, + *_rt_conv_scalars( + upper_pad=upper_pad_dhw, + lower_pad=lower_pad_dhw, + stride=stride_dhw, + dil=dil_dhw, + ), + current_stream, + ) + + workspace_count = 1 + if use_cold_l2: + one_workspace_bytes = ( + input_storage.numel() * input_storage.element_size() + + filter_storage.numel() * filter_storage.element_size() + + output_storage.numel() * output_storage.element_size() + ) + workspace_count = testing.get_workspace_count( + one_workspace_bytes, warmup_iterations, iterations + ) + + exec_time = testing.benchmark( + compiled_fn, + workspace_generator=generate_tensors, + workspace_count=workspace_count, + stream=current_stream, + warmup_iterations=warmup_iterations, + iterations=iterations, + use_cuda_graphs=True, + ) + runtime_s = exec_time / 1.0e6 + fmas = (N * Z * P * Q) * K * (C * T * R * S) + flop = 2 * fmas + gflop = flop / 1.0e9 + gflops = gflop / runtime_s + + print("Average Runtime : ", exec_time / 1000, "ms") + print("GFLOPS : ", gflops) + + return exec_time + + +# ///////////////////////////////////////////////////////////////////////////// +# Argument parsing and main +# ///////////////////////////////////////////////////////////////////////////// + + +def _parse_comma_separated_ints(s: str) -> Tuple[int, ...]: + try: + return tuple(int(x.strip()) for x in s.split(",")) + except ValueError: + raise argparse.ArgumentTypeError( + "Invalid format. Expected comma-separated integers." + ) + + +if __name__ == "__main__": + parser = argparse.ArgumentParser( + description="Blackwell GeForce (SM120) 3D convolution via implicit GEMM" + ) + + # Convolution parameters + parser.add_argument( + "--ncdhw", + type=_parse_comma_separated_ints, + default=(1, 64, 8, 8, 8), + help="Input tensor shape (N,C,D,H,W)", + ) + parser.add_argument( + "--ktrs", + type=_parse_comma_separated_ints, + default=(128, 3, 3, 3), + help="Filter tensor shape components (K,T,R,S)", + ) + parser.add_argument( + "--stride_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Stride (Sd,Sh,Sw)", + ) + parser.add_argument( + "--upper_pad_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Upper padding (PadD,PadH,PadW)", + ) + parser.add_argument( + "--lower_pad_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Lower padding (PadD,PadH,PadW)", + ) + parser.add_argument( + "--dil_dhw", + type=_parse_comma_separated_ints, + default=(1, 1, 1), + help="Dilation (DilD,DilH,DilW)", + ) + + # Data type parameters + parser.add_argument( + "--ab_dtype", + type=cutlass.dtype, + default=cutlass.Float16, + ) + parser.add_argument( + "--d_dtype", + type=cutlass.dtype, + default=cutlass.Float16, + ) + parser.add_argument( + "--acc_dtype", + type=cutlass.dtype, + default=cutlass.Float32, + ) + + # Tile shape + parser.add_argument( + "--tile_shape_mnk", + type=_parse_comma_separated_ints, + default=(128, 128, 32), + help="CTA tile shape (M,N,K)", + ) + + # Validation and benchmark parameters + parser.add_argument( + "--tolerance", type=float, default=1e-01, help="Tolerance for validation" + ) + parser.add_argument( + "--warmup_iterations", type=int, default=0, help="Warmup iterations" + ) + parser.add_argument( + "--iterations", + type=int, + default=1, + help="Number of iterations to run the kernel", + ) + parser.add_argument( + "--skip_ref_check", + action="store_true", + default=False, + help="Skip reference checking", + ) + parser.add_argument( + "--use_cold_l2", + action="store_true", + default=False, + help="Use circular buffer tensor sets to ensure L2 cold cache", + ) + + # Epilogue + parser.add_argument( + "--use_bias", + action="store_true", + default=False, + help="Add a per-output-channel bias, D = activation(acc + bias). Bias and " + "no-bias are separate cubins", + ) + parser.add_argument( + "--activation", + type=str, + default="identity", + choices=sorted(EPILOGUE_ACTIVATIONS), + help="Compile-time epilogue activation, D = activation(acc + bias). One " + "activation per cubin", + ) + + args = parser.parse_args() + + if len(args.ncdhw) != 5: + parser.error("--ncdhw must contain exactly 5 values (N,C,D,H,W)") + if len(args.ktrs) != 4: + parser.error("--ktrs must contain exactly 4 values (K,T,R,S)") + if len(args.tile_shape_mnk) != 3: + parser.error("--tile_shape_mnk must contain exactly 3 values (M,N,K)") + + run( + args.ncdhw, + args.ktrs, + args.stride_dhw, + args.upper_pad_dhw, + args.lower_pad_dhw, + args.dil_dhw, + args.ab_dtype, + args.d_dtype, + args.acc_dtype, + args.tile_shape_mnk, + args.tolerance, + args.warmup_iterations, + args.iterations, + args.use_cold_l2, + args.skip_ref_check, + args.use_bias, + args.activation, + ) + print("PASS") diff --git a/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/dense_gemm/dense_gemm.py b/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/dense_gemm/dense_gemm.py index 7fc7b42133..9ded72a616 100644 --- a/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/dense_gemm/dense_gemm.py +++ b/examples/python/CuTeDSL/cute/blackwell_geforce/kernel/dense_gemm/dense_gemm.py @@ -27,13 +27,14 @@ # OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. import argparse +import os from typing import Tuple, Type import cuda.bindings.driver as cuda import cutlass import cutlass.cute as cute -import cutlass.cute.testing as testing +from cutlass import testing import cutlass.utils as utils import cutlass.pipeline as pipeline import cutlass.utils.hopper_helpers as sm90_utils @@ -64,7 +65,7 @@ .. code-block:: bash - python examples/blackwell_geforce/dense_gemm.py \ + python examples/cute/blackwell_geforce/kernel/dense_gemm/dense_gemm.py \ --mnkl 8192,8192,8192,1 --tile_shape_mnk 128,256,64 \ --a_dtype Float16 --b_dtype Float16 \ --c_dtype Float16 --acc_dtype Float32 \ @@ -79,7 +80,7 @@ .. code-block:: bash - ncu python examples/blackwell_geforce/dense_gemm.py \ + ncu python examples/cute/blackwell_geforce/kernel/dense_gemm/dense_gemm.py \ --mnkl 8192,8192,8192,1 --tile_shape_mnk 128,256,64 \ --a_dtype Float16 --b_dtype Float16 \ --c_dtype Float16 --acc_dtype Float32 \ @@ -112,7 +113,9 @@ def parse_comma_separated_ints(s: str): def parse_arguments() -> argparse.Namespace: - parser = argparse.ArgumentParser(description="Example of MxNxKxL GEMM on Blackwell Geforce.") + parser = argparse.ArgumentParser( + description="Example of MxNxKxL GEMM on GeForce architectures." + ) parser.add_argument( "--mnkl", @@ -130,8 +133,9 @@ def parse_arguments() -> argparse.Namespace: (128, 128, 64), (128, 256, 64), (128, 128, 128), + (256, 128, 32), ], - default=(64, 64, 64), + default=(128, 128, 64), help="CTA tile shape (comma-separated)", ) parser.add_argument( @@ -157,6 +161,12 @@ def parse_arguments() -> argparse.Namespace: parser.add_argument("--a_major", choices=["k", "m"], type=str, default="k") parser.add_argument("--b_major", choices=["k", "n"], type=str, default="k") parser.add_argument("--c_major", choices=["n", "m"], type=str, default="n") + parser.add_argument( + "--epi_stage", + type=int, + default=4, + help="Number of epilogue shared-memory stages", + ) parser.add_argument( "--tolerance", type=float, default=1e-01, help="Tolerance for validation" ) @@ -187,6 +197,9 @@ def parse_arguments() -> argparse.Namespace: if len(args.mnkl) != 4: parser.error("--mnkl must contain exactly 4 values") + if args.epi_stage <= 0: + parser.error("--epi_stage must be greater than 0") + return args @@ -200,10 +213,12 @@ def __init__( self, acc_dtype, tile_shape_mnk, + epi_stage=4, ): self.acc_dtype = acc_dtype self.cluster_shape_mnk = (1, 1, 1) self.tile_shape_mnk = tuple(tile_shape_mnk) + self.epi_stage = epi_stage self.tiled_mma = None self.num_mcast_ctas_a = None self.num_mcast_ctas_b = None @@ -212,18 +227,18 @@ def __init__( self.occupancy = 1 # TODO: remove this hard code for user input ? - self.atom_layout = (2, 2, 1) + self.atom_layout = (4, 2, 1) self.num_mma_warps = ( self.atom_layout[0] * self.atom_layout[1] * self.atom_layout[2] ) self.num_threads_per_warp = 32 self.threads_per_cta = ( - self.num_mma_warps + 1 # 1 warp for DMA + self.num_mma_warps + 4 # 1 warp for DMA ) * self.num_threads_per_warp - self.smem_capacity = utils.get_smem_capacity_in_bytes("sm_120") + self.smem_capacity = cutlass.memory.get_smem_capacity_in_bytes() self.ab_stage = None - self.epi_stage = None + self.epi_stage = epi_stage self.a_smem_layout_staged = None self.b_smem_layout_staged = None @@ -281,10 +296,16 @@ def _setup_attributes(self): self.c_dtype, self.smem_capacity, self.occupancy, + self.epi_stage, ) import sys + print( + f"Computed pipeline stages: ab_stage={self.ab_stage}, " + f"epi_stage={self.epi_stage}" + ) + if self.ab_stage == 0: print("ab_stage == 0, no enough shared memory. This case will be skipped.") sys.exit(0) @@ -337,9 +358,9 @@ def __call__( self.b_dtype = b.element_type self.c_dtype = c.element_type - self.a_layout = utils.LayoutEnum.from_tensor(a) - self.b_layout = utils.LayoutEnum.from_tensor(b) - self.c_layout = utils.LayoutEnum.from_tensor(c) + self.a_layout = cutlass.tensor_utils.LayoutEnum.from_tensor(a) + self.b_layout = cutlass.tensor_utils.LayoutEnum.from_tensor(b) + self.c_layout = cutlass.tensor_utils.LayoutEnum.from_tensor(c) if cutlass.const_expr( self.a_dtype.width == 16 and self.a_dtype != self.b_dtype @@ -425,6 +446,8 @@ class SharedStorage: block=[self.threads_per_cta, 1, 1], cluster=[1, 1, 1], stream=stream, + # max_number_threads=[self.threads_per_cta, 1, 1], + min_blocks_per_mp=1, ) return @@ -515,7 +538,7 @@ def kernel( # ///////////////////////////////////////////////////////////////////////////// # Alloc and init AB full/empty + ACC full mbar (pipeline) # ///////////////////////////////////////////////////////////////////////////// - smem = cutlass.utils.SmemAllocator() + smem = cutlass.memory.SmemAllocator() storage = smem.allocate(self.shared_storage) # mbar arrays @@ -525,11 +548,8 @@ def kernel( mainloop_pipeline_producer_group = pipeline.CooperativeGroup( pipeline.Agent.Thread ) - # Each warp will constribute to the arrive count with the number of mcast size - mcast_size = self.num_mcast_ctas_a + self.num_mcast_ctas_b - 1 - consumer_arrive_cnt = mcast_size * self.num_mma_warps mainloop_pipeline_consumer_group = pipeline.CooperativeGroup( - pipeline.Agent.Thread, consumer_arrive_cnt + pipeline.Agent.Warp, self.num_mma_warps ) cta_layout_vmnk = cute.make_layout((1, *cta_layout_mnk.shape)) @@ -540,6 +560,7 @@ def kernel( tx_count=tma_copy_bytes, barrier_storage=mainloop_pipeline_array_ptr, cta_layout_vmnk=cta_layout_vmnk, + enable_multicast_signaling=True, ) # Cluster arrive after barrier init @@ -649,6 +670,36 @@ def kernel( num_k_blocks = cute.size(tCrA, mode=[2]) + # Monotonic sC buffer index across work tiles (wraps at epi_stage) + # so successive tiles write to distinct slots, not restarting at 0. + epi_buffer = cutlass.Int32(0) + + # Hoist PipelineTmaStore out of the work-tile loop so its lifetime + # spans the whole epilogue and pairs with the monotonic epi_buffer. + tma_store_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.num_mma_warps * self.num_threads_per_warp, + ) + tma_store_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.epi_stage, + producer_group=tma_store_producer_group, + ) + + # Monotonic sC buffer index across work tiles (wraps at epi_stage) + # so successive tiles write to distinct slots, not restarting at 0. + epi_buffer = cutlass.Int32(0) + + # Hoist PipelineTmaStore out of the work-tile loop so its lifetime + # spans the whole epilogue and pairs with the monotonic epi_buffer. + tma_store_producer_group = pipeline.CooperativeGroup( + pipeline.Agent.Thread, + self.num_mma_warps * self.num_threads_per_warp, + ) + tma_store_pipeline = pipeline.PipelineTmaStore.create( + num_stages=self.epi_stage, + producer_group=tma_store_producer_group, + ) + # /////////////////////////////////////////////////////////////////////////////// # Copy Atom A/B retiling for TMA load A/B # /////////////////////////////////////////////////////////////////////////////// @@ -716,6 +767,7 @@ def kernel( ) if k_block_idx == num_k_blocks - 1: + cute.arch.fence_view_async_shared() mainloop_pipeline.consumer_release(mainloop_consumer_state) mainloop_consumer_state.advance() @@ -762,6 +814,7 @@ def kernel( ) if k_block_idx == num_k_blocks - 1: + cute.arch.fence_view_async_shared() mainloop_pipeline.consumer_release(mainloop_consumer_state) mainloop_consumer_state.advance() @@ -816,11 +869,10 @@ def kernel( # (R2S, R2S_M, R2S_N) tRS_rAcc = tiled_copy_r2s.retile(accumulators) - # Allocate D registers. + # Allocate D registers (c_dtype, reused across all epi tiles). rD_shape = cute.shape(thr_copy_r2s.partition_S(sC)) tRS_rD_layout = cute.make_layout(rD_shape[:3]) - tRS_rD = cute.make_rmem_tensor(tRS_rD_layout.shape, self.acc_dtype) - size_tRS_rD = cute.size(tRS_rD) + tRS_rD_out = cute.make_rmem_tensor(tRS_rD_layout.shape, self.c_dtype) sepi_for_tma_partition = cute.group_modes(sC, 0, 2) tcgc_for_tma_partition = cute.zipped_divide(gC_mnl_slice, self.epi_tile) @@ -833,65 +885,68 @@ def kernel( tcgc_for_tma_partition, ) - epi_tile_num = cute.size(tcgc_for_tma_partition, mode=[1]) - epi_tile_shape = tcgc_for_tma_partition.shape[1] - epi_tile_layout = cute.make_layout( - epi_tile_shape, stride=(1, epi_tile_shape[0]) - ) - - # Initialize tma store pipeline - tma_store_producer_group = pipeline.CooperativeGroup( - pipeline.Agent.Thread, - self.num_mma_warps * self.num_threads_per_warp, - ) - tma_store_pipeline = pipeline.PipelineTmaStore.create( - num_stages=self.epi_stage, - producer_group=tma_store_producer_group, - ) + epi_rest_m = bSG_gD.shape[1][0] + epi_rest_n = bSG_gD.shape[1][1] + epi_tile_m = self.epi_tile[0] + epi_tile_n = self.epi_tile[1] + mma_tile_m = self.tile_shape_mnk[0] // cute.size(tRS_rAcc, mode=[1]) + mma_tile_n = self.tile_shape_mnk[1] // cute.size(tRS_rAcc, mode=[2]) + mma_m_per_epi_m = epi_tile_m // mma_tile_m + mma_n_per_epi_n = epi_tile_n // mma_tile_n + + for epi_m in cutlass.range_constexpr(epi_rest_m): + for epi_n in cutlass.range_constexpr(epi_rest_n): + for mma_n_in_epi in cutlass.range_constexpr(mma_n_per_epi_n): + for mma_m_in_epi in cutlass.range_constexpr( + mma_m_per_epi_m + ): + mma_m = epi_m * mma_m_per_epi_m + mma_m_in_epi + mma_n = epi_n * mma_n_per_epi_n + mma_n_in_epi + tRS_rD_slice = tRS_rD_out[ + (None, mma_m_in_epi, mma_n_in_epi) + ] + tRS_rAcc_slice = tRS_rAcc[(None, mma_m, mma_n)] + for elem_idx in cutlass.range_constexpr( + cute.size(tRS_rD_slice) + ): + tRS_rD_slice[elem_idx] = tRS_rAcc_slice[ + elem_idx + ].to(self.c_dtype) + + # Register to shared memory + cute.copy( + tiled_copy_r2s, + tRS_rD_out, + tRS_sD[(None, None, None, epi_buffer)], + ) - for epi_idx in cutlass.range_constexpr(epi_tile_num): - # Copy from accumulators to D registers - for epi_v in cutlass.range_constexpr(size_tRS_rD): - tRS_rD[epi_v] = tRS_rAcc[epi_idx * size_tRS_rD + epi_v] + cute.arch.fence_view_async_shared() + # barrier for sync + self.epilog_sync_barrier.arrive_and_wait() - # Type conversion - tRS_rD_out = cute.make_rmem_tensor( - tRS_rD_layout.shape, self.c_dtype - ) - acc_vec = tRS_rD.load() - tRS_rD_out.store(acc_vec.to(self.c_dtype)) - - # Register to shared memory - epi_buffer = epi_idx % cute.size(tRS_sD, mode=[3]) - cute.copy( - tiled_copy_r2s, - tRS_rD_out, - tRS_sD[(None, None, None, epi_buffer)], - ) + # Copy from shared memory to global memory + if warp_idx == 0: + cute.copy( + tma_atom_c, + bSG_sD[(None, epi_buffer)], + bSG_gD[(None, (epi_m, epi_n))], + ) + tma_store_pipeline.producer_commit() + tma_store_pipeline.producer_acquire() - cute.arch.fence_proxy( - "async.shared", - space="cta", - ) - # barrier for sync - self.epilog_sync_barrier.arrive_and_wait() + # Do not let non-store warps reuse sC until warp 0 has + # committed the TMA store and acquired the next store slot. + self.epilog_sync_barrier.arrive_and_wait() - # Get the global memory coordinate for the current epi tile. - gmem_coord = epi_tile_layout.get_hier_coord(epi_idx) - # Copy from shared memory to global memory - if warp_idx == 0: - cute.copy( - tma_atom_c, - bSG_sD[(None, epi_buffer)], - bSG_gD[(None, gmem_coord)], - ) - tma_store_pipeline.producer_commit() - tma_store_pipeline.producer_acquire() + # Advance the monotonic buffer counter; wraps at epi_stage. + epi_buffer = (epi_buffer + 1) % self.epi_stage # Advance to the next work tile tile_sched.advance_to_next_work() work_tile = tile_sched.get_current_work() - tma_store_pipeline.producer_tail() + if warp_idx == 0: + tma_store_pipeline.producer_tail() + self.epilog_sync_barrier.arrive_and_wait() # End of for k_tile loop # End of while loop # End of MMA warp group @@ -953,6 +1008,8 @@ def kernel( # Wait A/B buffer empty mainloop_pipeline.producer_tail(mainloop_producer_state) + else: + cute.arch.setmaxregister_decrease(self.load_register_requirement) return @staticmethod @@ -964,6 +1021,7 @@ def _compute_stages( c_dtype: type[cutlass.Numeric], smem_capacity: int, occupancy: int, + epi_stage: int = 4, ) -> tuple[int, int]: """Computes the number of stages for A/B/C operands based on heuristics. @@ -982,23 +1040,19 @@ def _compute_stages( (A/B operand stages, epilogue stages) :rtype: tuple[int, int] """ - epi_stage = 8 - c_bytes_per_stage = cute.size(epi_tile) * c_dtype.width // 8 - epi_bytes = c_bytes_per_stage * epi_stage - - a_shape = cute.slice_(tile_shape_mnk, (None, 0, None)) - b_shape = cute.slice_(tile_shape_mnk, (0, None, None)) + c_bytes_per_stage = epi_tile[0] * epi_tile[1] * c_dtype.width // 8 ab_bytes_per_stage = ( - cute.size(a_shape) * a_dtype.width // 8 - + cute.size(b_shape) * b_dtype.width // 8 + tile_shape_mnk[0] * tile_shape_mnk[2] * a_dtype.width // 8 + + tile_shape_mnk[1] * tile_shape_mnk[2] * b_dtype.width // 8 ) mbar_helpers_bytes = 1024 - ab_stage = ( - (smem_capacity - occupancy * 1024) // occupancy - - mbar_helpers_bytes - - epi_bytes - ) // ab_bytes_per_stage + smem_budget_bytes = ( + smem_capacity - occupancy * 1024 + ) // occupancy - mbar_helpers_bytes + epi_bytes = c_bytes_per_stage * epi_stage + + ab_stage = (smem_budget_bytes - epi_bytes) // ab_bytes_per_stage return ab_stage, epi_stage @staticmethod @@ -1173,12 +1227,14 @@ def run( iterations: int, skip_ref_check: bool, use_cold_l2: bool = False, + epi_stage: int = 4, **kwargs, ): import torch import cutlass.torch as cutlass_torch - print("Running Blackwell Geforce Dense GEMM with:") + print("Running GeForce Dense GEMM with:") + print(f"Target arch (CUTE_DSL_ARCH): {os.environ.get('CUTE_DSL_ARCH', 'not set')}") print(f"mnkl: {mnkl}") print( f"A dtype: {a_dtype}, B dtype: {b_dtype}, C dtype: {c_dtype}, Acc dtype: {acc_dtype}" @@ -1190,6 +1246,7 @@ def run( print(f"Iterations: {iterations}") print(f"Skip reference checking: {skip_ref_check}") print(f"Use cold L2: {use_cold_l2}") + print(f"Epi stage: {epi_stage}") a_dtype = getattr(cutlass, a_dtype) if isinstance(a_dtype, str) else a_dtype b_dtype = getattr(cutlass, b_dtype) if isinstance(b_dtype, str) else b_dtype @@ -1204,9 +1261,9 @@ def run( if not torch.cuda.is_available(): raise RuntimeError("GPU is required to run this example!") - a_torch_cpu = cutlass_torch.matrix(l, m, k, a_major, a_dtype) - b_torch_cpu = cutlass_torch.matrix(l, n, k, b_major, b_dtype) - c_torch_cpu = cutlass_torch.matrix(l, m, n, c_major, c_dtype) + a_torch_cpu = cutlass_torch.matrix(l, m, k, a_major == "m", a_dtype) + b_torch_cpu = cutlass_torch.matrix(l, n, k, b_major == "n", b_dtype) + c_torch_cpu = cutlass_torch.matrix(l, m, n, c_major == "m", c_dtype) def create_cute_tensor(data_ref, cutlass_dtype): cute_tensor, torch_tensor = cutlass_torch.cute_tensor_like( @@ -1230,6 +1287,7 @@ def create_cute_tensor(data_ref, cutlass_dtype): gemm = Sm120GemmKernel( acc_dtype, tile_shape_mnk, + epi_stage, ) # Compute max active clusters on current device @@ -1271,9 +1329,9 @@ def create_cute_tensor(data_ref, cutlass_dtype): ) def generate_tensors(): - a_torch_cpu = cutlass_torch.matrix(l, m, k, a_major, a_dtype) - b_torch_cpu = cutlass_torch.matrix(l, n, k, b_major, b_dtype) - c_torch_cpu = cutlass_torch.matrix(l, m, n, c_major, c_dtype) + a_torch_cpu = cutlass_torch.matrix(l, m, k, a_major == "m", a_dtype) + b_torch_cpu = cutlass_torch.matrix(l, n, k, b_major == "n", b_dtype) + c_torch_cpu = cutlass_torch.matrix(l, m, n, c_major == "m", c_dtype) mA_workspace, _ = create_cute_tensor(a_torch_cpu, a_dtype) mB_workspace, _ = create_cute_tensor(b_torch_cpu, b_dtype) mC_workspace, _ = create_cute_tensor(c_torch_cpu, c_dtype) @@ -1321,5 +1379,6 @@ def generate_tensors(): args.iterations, args.skip_ref_check, args.use_cold_l2, + args.epi_stage, ) print("PASS") diff --git a/examples/python/CuTeDSL/cute/notebooks/data_types.ipynb b/examples/python/CuTeDSL/cute/notebooks/data_types.ipynb deleted file mode 100644 index dd05177a71..0000000000 --- a/examples/python/CuTeDSL/cute/notebooks/data_types.ipynb +++ /dev/null @@ -1,270 +0,0 @@ -{ - "cells": [ - { - "cell_type": "code", - "execution_count": 1, - "metadata": {}, - "outputs": [], - "source": [ - "import cutlass\n", - "import cutlass.cute as cute" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Understanding data structure in CuTe DSL\n", - "\n", - "In most cases, data structures in CuTe DSL work the same as Python data structures with the notable difference that Python data structures in most cases are considered as static data which are interpreted by the DSL compiler embedded inside Python interpreter.\n", - "\n", - "To differentiate between compile-time and runtime values, CuTe DSL introduces primitive types that \n", - "represent dynamic values in JIT-compiled code.\n", - "\n", - "CuTe DSL provides a comprehensive set of primitive numeric types for representing dynamic values at \n", - "runtime. These types are formally defined within the CuTe DSL typing system:\n", - "\n", - "### Integer Types\n", - "- `Int8` - 8-bit signed integer\n", - "- `Int16` - 16-bit signed integer \n", - "- `Int32` - 32-bit signed integer\n", - "- `Int64` - 64-bit signed integer\n", - "- `Int128` - 128-bit signed integer\n", - "- `Uint8` - 8-bit unsigned integer\n", - "- `Uint16` - 16-bit unsigned integer\n", - "- `Uint32` - 32-bit unsigned integer\n", - "- `Uint64` - 64-bit unsigned integer\n", - "- `Uint128` - 128-bit unsigned integer\n", - "\n", - "### Floating Point Types\n", - "- `Float16` - 16-bit floating point\n", - "- `Float32` - 32-bit floating point \n", - "- `Float64` - 64-bit floating point\n", - "- `BFloat16` - Brain Floating Point format (16-bit)\n", - "- `TFloat32` - Tensor Float32 format (reduced precision format used in tensor operations)\n", - "- `Float8E4M3` - 8-bit floating point with 4-bit exponent and 3-bit mantissa\n", - "- `Float8E5M2` - 8-bit floating point with 5-bit exponent and 2-bit mantissa\n", - "\n", - "These specialized types are designed to represent dynamic values in CuTe DSL code that will be \n", - "evaluated at runtime, in contrast to Python's built-in numeric types which are evaluated during \n", - "compilation.\n", - "\n", - "### Example usage:\n", - "\n", - "```python\n", - "x = cutlass.Int32(5) # Creates a 32-bit integer\n", - "y = cutlass.Float32(3.14) # Creates a 32-bit float\n", - "\n", - "@cute.jit\n", - "def foo(a: cutlass.Int32): # annotate `a` as 32-bit integer passed to jit function via ABI\n", - " ...\n", - "```\n" - ] - }, - { - "cell_type": "code", - "execution_count": 2, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "a(static) = ?\n", - "b(static) = ?\n", - "a(dynamic) = 3.140000\n", - "b(dynamic) = 5\n" - ] - } - ], - "source": [ - "@cute.jit\n", - "def bar():\n", - " a = cutlass.Float32(3.14)\n", - " print(\"a(static) =\", a) # prints `a(static) = ?`\n", - " cute.printf(\"a(dynamic) = {}\", a) # prints `a(dynamic) = 3.140000`\n", - "\n", - " b = cutlass.Int32(5)\n", - " print(\"b(static) =\", b) # prints `b(static) = 5`\n", - " cute.printf(\"b(dynamic) = {}\", b) # prints `b(dynamic) = 5`\n", - "\n", - "\n", - "bar()" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "### Type Conversion API\n", - "\n", - "CUTLASS numeric types provide type conversion through the `to()` method available on all Numeric types. This allows you to convert between different numeric data types at runtime.\n", - "\n", - "Syntax:\n", - "\n", - "```python\n", - "new_value = value.to(target_type)\n", - "```\n", - "\n", - "The `to()` method supports conversion between:\n", - "- Integer types (Int8, Int16, Int32, Int64, UInt8, UInt16, UInt32, UInt64)\n", - "- Floating point types (Float16, Float32, Float64, BFloat16)\n", - "- Mixed integer/floating point conversions\n", - "\n", - "Note that when converting from floating point to integer types, the decimal portion is truncated. When converting between types with different ranges, values may be clamped or lose precision if they exceed the target type's representable range." - ] - }, - { - "cell_type": "code", - "execution_count": 3, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Int32(42) => Float32(42.000000)\n", - "Float32(3.140000) => Int32(3)\n", - "Int32(127) => Int8(127)\n", - "Int32(300) => Int8(44) (truncated due to range limitation)\n" - ] - } - ], - "source": [ - "@cute.jit\n", - "def type_conversion():\n", - " # Convert from Int32 to Float32\n", - " x = cutlass.Int32(42)\n", - " y = x.to(cutlass.Float32)\n", - " cute.printf(\"Int32({}) => Float32({})\", x, y)\n", - "\n", - " # Convert from Float32 to Int32\n", - " a = cutlass.Float32(3.14)\n", - " b = a.to(cutlass.Int32)\n", - " cute.printf(\"Float32({}) => Int32({})\", a, b)\n", - "\n", - " # Convert from Int32 to Int8\n", - " c = cutlass.Int32(127)\n", - " d = c.to(cutlass.Int8)\n", - " cute.printf(\"Int32({}) => Int8({})\", c, d)\n", - "\n", - " # Convert from Int32 to Int8 with value exceeding Int8 range\n", - " e = cutlass.Int32(300)\n", - " f = e.to(cutlass.Int8)\n", - " cute.printf(\"Int32({}) => Int8({}) (truncated due to range limitation)\", e, f)\n", - "\n", - "\n", - "type_conversion()" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "### Operator Overloading\n", - "\n", - "CUTLASS numeric types support Python's built-in operators, allowing you to write natural mathematical expressions. The operators work with both CUTLASS numeric types and Python native numeric types.\n", - "\n", - "Supported operators include:\n", - "- Arithmetic: `+`, `-`, `*`, `/`, `//`, `%`, `**`\n", - "- Comparison: `<`, `<=`, `==`, `!=`, `>=`, `>`\n", - "- Bitwise: `&`, `|`, `^`, `<<`, `>>`\n", - "- Unary: `-` (negation), `~` (bitwise NOT)" - ] - }, - { - "cell_type": "code", - "execution_count": 4, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "a: Int32(10), b: Int32(3)\n", - "x: Float32(5.500000)\n", - "\n", - "a + b = 13\n", - "x * 2 = 11.000000\n", - "a + x = 15.500000 (Int32 + Float32 promotes to Float32)\n", - "a / b = 3.333333\n", - "x / 2.0 = 2.750000\n", - "a > b = 1\n", - "a & b = 2\n", - "-a = -10\n", - "~a = -11\n" - ] - } - ], - "source": [ - "@cute.jit\n", - "def operator_demo():\n", - " # Arithmetic operators\n", - " a = cutlass.Int32(10)\n", - " b = cutlass.Int32(3)\n", - " cute.printf(\"a: Int32({}), b: Int32({})\", a, b)\n", - "\n", - " x = cutlass.Float32(5.5)\n", - " cute.printf(\"x: Float32({})\", x)\n", - "\n", - " cute.printf(\"\")\n", - "\n", - " sum_result = a + b\n", - " cute.printf(\"a + b = {}\", sum_result)\n", - "\n", - " y = x * 2 # Multiplying with Python native type\n", - " cute.printf(\"x * 2 = {}\", y)\n", - "\n", - " # Mixed type arithmetic (Int32 + Float32) that integer is converted into float32\n", - " mixed_result = a + x\n", - " cute.printf(\"a + x = {} (Int32 + Float32 promotes to Float32)\", mixed_result)\n", - "\n", - " # Division with Int32 (note: integer division)\n", - " div_result = a / b\n", - " cute.printf(\"a / b = {}\", div_result)\n", - "\n", - " # Float division\n", - " float_div = x / cutlass.Float32(2.0)\n", - " cute.printf(\"x / 2.0 = {}\", float_div)\n", - "\n", - " # Comparison operators\n", - " is_greater = a > b\n", - " cute.printf(\"a > b = {}\", is_greater)\n", - "\n", - " # Bitwise operators\n", - " bit_and = a & b\n", - " cute.printf(\"a & b = {}\", bit_and)\n", - "\n", - " neg_a = -a\n", - " cute.printf(\"-a = {}\", neg_a)\n", - "\n", - " not_a = ~a\n", - " cute.printf(\"~a = {}\", not_a)\n", - "\n", - "\n", - "operator_demo()" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": "Python 3 (ipykernel)", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.12.5" - } - }, - "nbformat": 4, - "nbformat_minor": 4 -} diff --git a/examples/python/CuTeDSL/cute/notebooks/hello_world.ipynb b/examples/python/CuTeDSL/cute/notebooks/hello_world.ipynb deleted file mode 100644 index 6bf35b76a9..0000000000 --- a/examples/python/CuTeDSL/cute/notebooks/hello_world.ipynb +++ /dev/null @@ -1,181 +0,0 @@ -{ - "cells": [ - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "# Your First Program with CuTe DSL\n", - "\n", - "## Introduction\n", - "\n", - "Welcome! In this tutorial, we'll write a simple \"Hello World\" program that runs on your GPU using CuTe DSL. This will help you understand the basics of GPU programming with our framework.\n", - "\n", - "### What You'll Learn\n", - "\n", - "- How to write code that runs on both CPU (host) and GPU (device),\n", - "- How to launch a GPU kernel (a function that runs on the GPU),\n", - "- Basic CUDA concepts like threads and thread blocks,\n", - "\n", - "### Step 1: Import Required Libraries\n", - "\n", - "First, let's import the libraries we need:" - ] - }, - { - "cell_type": "code", - "execution_count": 1, - "metadata": {}, - "outputs": [], - "source": [ - "import cutlass\n", - "import cutlass.cute as cute" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "\n", - "### Step 2: Write Our GPU Kernel\n", - "A GPU kernel is a function that runs on the GPU. Here's a simple kernel that prints \"Hello World\".\n", - "Key concepts:\n", - "- `@cute.kernel`: This decorator tells CUTLASS that this function should run on the GPU\n", - "- `cute.arch.thread_idx()`: Gets the ID of the current GPU thread (like a worker's ID number)\n", - "- We only want one thread to print the message (thread 0) to avoid multiple prints" - ] - }, - { - "cell_type": "code", - "execution_count": 2, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.kernel\n", - "def kernel():\n", - " # Get the x component of the thread index (y and z components are unused)\n", - " tidx, _, _ = cute.arch.thread_idx()\n", - " # Only the first thread (thread 0) prints the message\n", - " if cutlass.dynamic_expr(tidx == 0):\n", - " cute.printf(\"Hello world\")" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "### Step 3: Write Our Host Function\n", - "\n", - "Now we need a function that sets up the GPU and launches our kernel.\n", - "Key concepts:\n", - "- `@cute.jit`: This decorator is for functions that run on the CPU but can launch GPU code\n", - "- We need to initialize CUDA before using the GPU\n", - "- `.launch()` tells CUDA how many blocks, threads, shared memory, etc. to use" - ] - }, - { - "cell_type": "code", - "execution_count": 3, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def hello_world():\n", - " # Print hello world from host code\n", - " cute.printf(\"hello world\")\n", - "\n", - " # Launch kernel\n", - " kernel().launch(\n", - " grid=(1, 1, 1), # Single thread block\n", - " block=(32, 1, 1), # One warp (32 threads) per thread block\n", - " )" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "### Step 4: Run Our Program\n", - "\n", - "There are 2 ways we can run our program:\n", - "\n", - "1. compile and run immediately\n", - "2. separate compilation which allows us to compile the code once and run multiple times\n", - " \n", - "Please note the `Compiling...` for Method 2 prints before the \"Hello world\" of the first kernel. This shows the asynchronous behavior between CPU and GPU prints. " - ] - }, - { - "cell_type": "code", - "execution_count": 5, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Running hello_world()...\n", - "Compiling...\n", - "hello world\n", - "Hello world\n", - "Compiling with PTX/CUBIN dumped...\n", - "Running compiled version...\n", - "hello world\n", - "Hello world\n" - ] - } - ], - "source": [ - "# Initialize CUDA context for launching a kernel with error checking\n", - "# We make context initialization explicit to allow users to control the context creation\n", - "# and avoid potential issues with multiple contexts\n", - "cutlass.cuda.initialize_cuda_context()\n", - "\n", - "# Method 1: Just-In-Time (JIT) compilation - compiles and runs the code immediately\n", - "print(\"Running hello_world()...\")\n", - "hello_world()\n", - "\n", - "# Method 2: Compile first (useful if you want to run the same code multiple times)\n", - "print(\"Compiling...\")\n", - "hello_world_compiled = cute.compile(hello_world)\n", - "\n", - "# Dump PTX/CUBIN files while compiling\n", - "from cutlass.cute import KeepPTX, KeepCUBIN\n", - "\n", - "print(\"Compiling with PTX/CUBIN dumped...\")\n", - "hello_world_compiled_ptx_on = cute.compile[KeepPTX, KeepCUBIN](hello_world)\n", - "\n", - "# Run the pre-compiled version\n", - "print(\"Running compiled version...\")\n", - "hello_world_compiled()" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": "Python 3 (ipykernel)", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.12.5" - }, - "widgets": { - "application/vnd.jupyter.widget-state+json": { - "state": {}, - "version_major": 2, - "version_minor": 0 - } - } - }, - "nbformat": 4, - "nbformat_minor": 4 -} diff --git a/examples/python/CuTeDSL/cute/notebooks/print.ipynb b/examples/python/CuTeDSL/cute/notebooks/print.ipynb deleted file mode 100644 index 0e70d2c43b..0000000000 --- a/examples/python/CuTeDSL/cute/notebooks/print.ipynb +++ /dev/null @@ -1,500 +0,0 @@ -{ - "cells": [ - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "# Printing with CuTe DSL\n", - "\n", - "This notebook demonstrates the different ways to print values in CuTe and explains the important distinction between static (compile-time) and dynamic (runtime) values.\n", - "\n", - "## Key Concepts\n", - "- Static values: Known at compile time\n", - "- Dynamic values: Only known at runtime\n", - "- Different printing methods for different scenarios\n", - "- Layout representation in CuTe\n", - "- Tensor visualization and formatting" - ] - }, - { - "cell_type": "code", - "execution_count": 1, - "metadata": {}, - "outputs": [], - "source": [ - "import cutlass\n", - "import cutlass.cute as cute\n", - "import numpy as np" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Print Example Function\n", - "\n", - "The `print_example` function demonstrates several important concepts:\n", - "\n", - "### 1. Python's `print` vs CuTe's `cute.printf`\n", - "- `print`: Can only show static values at compile time\n", - "- `cute.printf`: Can display both static and dynamic values at runtime\n", - "\n", - "### 2. Value Types\n", - "- `a`: Dynamic `Int32` value (runtime)\n", - "- `b`: Static `Constexpr[int]` value (compile-time)\n", - "\n", - "### 3. Layout Printing\n", - "Shows how layouts are represented differently in static vs dynamic contexts:\n", - "- Static context: Unknown values shown as `?`\n", - "- Dynamic context: Actual values displayed" - ] - }, - { - "cell_type": "code", - "execution_count": 2, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def print_example(a: cutlass.Int32, b: cutlass.Constexpr[int]):\n", - " \"\"\"\n", - " Demonstrates different printing methods in CuTe and how they handle static vs dynamic values.\n", - "\n", - " This example shows:\n", - " 1. How Python's `print` function works with static values at compile time but can't show dynamic values\n", - " 2. How `cute.printf` can display both static and dynamic values at runtime\n", - " 3. The difference between types in static vs dynamic contexts\n", - " 4. How layouts are represented in both printing methods\n", - "\n", - " Args:\n", - " a: A dynamic Int32 value that will be determined at runtime\n", - " b: A static (compile-time constant) integer value\n", - " \"\"\"\n", - " # Use Python `print` to print static information\n", - " print(\">>>\", b) # => 2\n", - " # `a` is dynamic value\n", - " print(\">>>\", a) # => ?\n", - "\n", - " # Use `cute.printf` to print dynamic information\n", - " cute.printf(\">?? {}\", a) # => 8\n", - " cute.printf(\">?? {}\", b) # => 2\n", - "\n", - " print(\">>>\", type(a)) # => \n", - " print(\">>>\", type(b)) # => \n", - "\n", - " layout = cute.make_layout((a, b))\n", - " print(\">>>\", layout) # => (?,2):(1,?)\n", - " cute.printf(\">?? {}\", layout) # => (8,2):(1,8)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Compile and Run\n", - "\n", - "**Direct Compilation and Run**\n", - " - `print_example(cutlass.Int32(8), 2)`\n", - " - Compiles and runs in one step will execute both static and dynamic print\n", - " * `>>>` stands for static print\n", - " * `>??` stands for dynamic print" - ] - }, - { - "cell_type": "code", - "execution_count": 3, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - ">>> 2\n", - ">>> ?\n", - ">>> Int32\n", - ">>> \n", - ">>> (?,2):(1,?)\n", - ">?? 8\n", - ">?? 2\n", - ">?? (8,2):(1,8)\n" - ] - } - ], - "source": [ - "print_example(cutlass.Int32(8), 2)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Compile Function\n", - "\n", - "When compiles the function with `cute.compile(print_example, cutlass.Int32(8), 2)`, Python interpreter \n", - "traces code and only evaluate static expression and print static information." - ] - }, - { - "cell_type": "code", - "execution_count": 4, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - ">>> 2\n", - ">>> ?\n", - ">>> Int32\n", - ">>> \n", - ">>> (?,2):(1,?)\n" - ] - } - ], - "source": [ - "print_example_compiled = cute.compile(print_example, cutlass.Int32(8), 2)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Call compiled function\n", - "\n", - "Only print out runtime information" - ] - }, - { - "cell_type": "code", - "execution_count": 5, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - ">?? 8\n", - ">?? 2\n", - ">?? (8,2):(1,8)\n" - ] - } - ], - "source": [ - "print_example_compiled(cutlass.Int32(8))" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Format String Example\n", - "\n", - "The `format_string_example` function shows an important limitation:\n", - "- F-strings in CuTe are evaluated at compile time\n", - "- This means dynamic values won't show their runtime values in f-strings\n", - "- Use `cute.printf` when you need to see runtime values" - ] - }, - { - "cell_type": "code", - "execution_count": 6, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Direct run output:\n", - "a: ?, b: 2\n", - "layout: (?,2):(1,?)\n" - ] - } - ], - "source": [ - "@cute.jit\n", - "def format_string_example(a: cutlass.Int32, b: cutlass.Constexpr[int]):\n", - " \"\"\"\n", - " Format string is evaluated at compile time.\n", - " \"\"\"\n", - " print(f\"a: {a}, b: {b}\")\n", - "\n", - " layout = cute.make_layout((a, b))\n", - " print(f\"layout: {layout}\")\n", - "\n", - "\n", - "print(\"Direct run output:\")\n", - "format_string_example(cutlass.Int32(8), 2)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Printing Tensor Examples\n", - "\n", - "CuTe provides specialized functionality for printing tensors through the `print_tensor` operation. The `cute.print_tensor` takes the following parameter:\n", - "- `Tensor` (required): A CuTe tensor object that you want to print. The tensor must support load and store operations\n", - "- `verbose` (optional, default=False): A boolean flag that controls the level of detail in the output. When set to True, it will print indices details for each element in the tensor.\n", - "\n", - "Below example code shows the difference between verbose ON and OFF, and how to print a sub range of the given tensor." - ] - }, - { - "cell_type": "code", - "execution_count": 7, - "metadata": {}, - "outputs": [], - "source": [ - "from cutlass.cute.runtime import from_dlpack\n", - "\n", - "\n", - "@cute.jit\n", - "def print_tensor_basic(x: cute.Tensor):\n", - " # Print the tensor\n", - " print(\"Basic output:\")\n", - " cute.print_tensor(x)\n", - "\n", - "\n", - "@cute.jit\n", - "def print_tensor_verbose(x: cute.Tensor):\n", - " # Print the tensor with verbose mode\n", - " print(\"Verbose output:\")\n", - " cute.print_tensor(x, verbose=True)\n", - "\n", - "\n", - "@cute.jit\n", - "def print_tensor_slice(x: cute.Tensor, coord: tuple):\n", - " # slice a 2D tensor from the 3D tensor\n", - " sliced_data = cute.slice_(x, coord)\n", - " y = cute.make_rmem_tensor(sliced_data.layout, sliced_data.element_type)\n", - " # Convert to TensorSSA format by loading the sliced data into the fragment\n", - " y.store(sliced_data.load())\n", - " print(\"Slice output:\")\n", - " cute.print_tensor(y)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "The default `cute.print_tensor` will output CuTe tensor with datatype, storage space, CuTe layout information, and print data in torch-style format." - ] - }, - { - "cell_type": "code", - "execution_count": 8, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Basic output:\n", - "tensor(raw_ptr(0x000000000a5f1d50: f32, generic, align<4>) o (4,3,2):(6,2,1), data=\n", - " [[[ 0.000000, 2.000000, 4.000000, ],\n", - " [ 6.000000, 8.000000, 10.000000, ],\n", - " [ 12.000000, 14.000000, 16.000000, ],\n", - " [ 18.000000, 20.000000, 22.000000, ]],\n", - "\n", - " [[ 1.000000, 3.000000, 5.000000, ],\n", - " [ 7.000000, 9.000000, 11.000000, ],\n", - " [ 13.000000, 15.000000, 17.000000, ],\n", - " [ 19.000000, 21.000000, 23.000000, ]]])\n" - ] - } - ], - "source": [ - "def tensor_print_example1():\n", - " shape = (4, 3, 2)\n", - "\n", - " # Creates [0,...,23] and reshape to (4, 3, 2)\n", - " data = np.arange(24, dtype=np.float32).reshape(*shape)\n", - "\n", - " print_tensor_basic(from_dlpack(data))\n", - "\n", - "\n", - "tensor_print_example1()" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "The verbosed print will show coodination details of each element in the tensor. The below example shows how we index element in a 2D 4x3 tensor space." - ] - }, - { - "cell_type": "code", - "execution_count": 9, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Verbose output:\n", - "tensor(raw_ptr(0x000000000a814cc0: f32, generic, align<4>) o (4,3):(3,1), data= (\n", - "\t(0,0)= 0.000000\n", - "\t(0,1)= 1.000000\n", - "\t(0,2)= 2.000000\n", - "\t(1,0)= 3.000000\n", - "\t(1,1)= 4.000000\n", - "\t(1,2)= 5.000000\n", - "\t(2,0)= 6.000000\n", - "\t(2,1)= 7.000000\n", - "\t(2,2)= 8.000000\n", - "\t(3,0)= 9.000000\n", - "\t(3,1)= 10.000000\n", - "\t(3,2)= 11.000000\n", - ")\n" - ] - } - ], - "source": [ - "def tensor_print_example2():\n", - " shape = (4, 3)\n", - "\n", - " # Creates [0,...,11] and reshape to (4, 3)\n", - " data = np.arange(12, dtype=np.float32).reshape(*shape)\n", - "\n", - " print_tensor_verbose(from_dlpack(data))\n", - "\n", - "\n", - "tensor_print_example2()" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "To print a subset elements in the given Tensor, we can use cute.slice_ to select a range of the given tensor, load them into register and then print the values with `cute.print_tensor`." - ] - }, - { - "cell_type": "code", - "execution_count": 10, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "Slice output:\n", - "tensor(raw_ptr(0x00007ffeeae1fc60: f32, rmem, align<32>) o (4):(3), data=\n", - " [ 0.000000, ],\n", - " [ 3.000000, ],\n", - " [Slice output:\n", - " 6.000000, ],\n", - " [ 9.000000, ])\n", - "tensor(raw_ptr(0x00007ffeeae1fc60: f32, rmem, align<32>) o (3):(1), data=\n", - " [ 3.000000, ],\n", - " [ 4.000000, ],\n", - " [ 5.000000, ])\n" - ] - } - ], - "source": [ - "def tensor_print_example3():\n", - " shape = (4, 3)\n", - "\n", - " # Creates [0,...,11] and reshape to (4, 3)\n", - " data = np.arange(12, dtype=np.float32).reshape(*shape)\n", - "\n", - " print_tensor_slice(from_dlpack(data), (None, 0))\n", - " print_tensor_slice(from_dlpack(data), (1, None))\n", - "\n", - "\n", - "tensor_print_example3()" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "To print the tensor in device memory, you can use `cute.print_tensor` within CuTe JIT kernels." - ] - }, - { - "cell_type": "code", - "execution_count": 13, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.kernel\n", - "def print_tensor_gpu(src: cute.Tensor):\n", - " print(src)\n", - " cute.print_tensor(src)\n", - "\n", - "\n", - "@cute.jit\n", - "def print_tensor_host(src: cute.Tensor):\n", - " print_tensor_gpu(src).launch(grid=(1, 1, 1), block=(1, 1, 1))" - ] - }, - { - "cell_type": "code", - "execution_count": 15, - "metadata": {}, - "outputs": [ - { - "name": "stdout", - "output_type": "stream", - "text": [ - "tensor o (4,3):(3,1)>\n" - ] - }, - { - "name": "stdout", - "output_type": "stream", - "text": [ - "tensor(raw_ptr(0x00007f5f81200400: f32, gmem, align<4>) o (4,3):(3,1), data=\n", - " [[-0.690547, -0.274619, -1.659539, ],\n", - " [-1.843524, -1.648711, 1.163431, ],\n", - " [-0.716668, -1.900705, 0.592515, ],\n", - " [ 0.711333, -0.552422, 0.860237, ]])\n" - ] - } - ], - "source": [ - "import torch\n", - "\n", - "\n", - "def tensor_print_example4():\n", - " a = torch.randn(4, 3, device=\"cuda\")\n", - " cutlass.cuda.initialize_cuda_context()\n", - " print_tensor_host(from_dlpack(a))\n", - "\n", - "\n", - "tensor_print_example4()" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "Currently, `cute.print_tensor` only supports tensor with integer data types and `Float16`/`Float32`/`Float64` floating point data types. We will support more data types in the future." - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": "Python 3 (ipykernel)", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.12.5" - } - }, - "nbformat": 4, - "nbformat_minor": 4 -} diff --git a/examples/python/CuTeDSL/cute/notebooks/tensor.ipynb b/examples/python/CuTeDSL/cute/notebooks/tensor.ipynb deleted file mode 100644 index 6106aae3b0..0000000000 --- a/examples/python/CuTeDSL/cute/notebooks/tensor.ipynb +++ /dev/null @@ -1,330 +0,0 @@ -{ - "cells": [ - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "import cutlass\n", - "import cutlass.cute as cute" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Tensor\n", - "\n", - "A tensor in CuTe is created through the composition of two key components:\n", - "\n", - "1. An **Engine** (E) - A random-access, pointer-like object that supports:\n", - " - Offset operation: `e + d → e` (offset engine by elements of a layout's codomain)\n", - " - Dereference operation: `*e → v` (dereference engine to produce value)\n", - "\n", - "2. A **Layout** (L) - Defines the mapping from coordinates to offsets\n", - "\n", - "A tensor is formally defined as the composition of an engine E with a layout L, expressed as `T = E ∘ L`. When evaluating a tensor at coordinate c, it:\n", - "\n", - "1. Maps the coordinate c to the codomain using the layout\n", - "2. Offsets the engine accordingly\n", - "3. Dereferences the result to obtain the tensor's value\n", - "\n", - "This can be expressed mathematically as:\n", - "\n", - "```\n", - "T(c) = (E ∘ L)(c) = *(E + L(c))\n", - "```\n", - "\n", - "## Example Usage\n", - "\n", - "Here's a simple example of creating a tensor using pointer and layout `(8,5):(5,1)` and fill with ones:" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def create_tensor_from_ptr(ptr: cute.Pointer):\n", - " layout = cute.make_layout((8, 5), stride=(5, 1))\n", - " tensor = cute.make_tensor(ptr, layout)\n", - " tensor.fill(1)\n", - " cute.print_tensor(tensor)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "This creates a tensor where:\n", - "- The engine is a pointer\n", - "- The layout with shape `(8, 5)` and stride `(5, 1)`\n", - "- The resulting tensor can be evaluated using coordinates defined by the layout\n", - "\n", - "We can test this by allocating buffer with torch and run test with pointer to torch tensor" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "import torch\n", - "\n", - "from cutlass.torch import dtype as torch_dtype\n", - "import cutlass.cute.runtime as cute_rt\n", - "\n", - "a = torch.randn(8, 5, dtype=torch_dtype(cutlass.Float32))\n", - "ptr_a = cute_rt.make_ptr(cutlass.Float32, a.data_ptr())\n", - "\n", - "create_tensor_from_ptr(ptr_a)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## DLPACK support \n", - "\n", - "CuTe DSL is designed to support dlpack protocol natively. This offers easy integration with frameworks \n", - "supporting DLPack, e.g. torch, numpy, jax, tensorflow, etc.\n", - "\n", - "For more information, please refer to DLPACK project: https://github.com/dmlc/dlpack\n", - "\n", - "Calling `from_dlpack` can convert any tensor or ndarray object supporting `__dlpack__` and `__dlpack_device__`.\n" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "from cutlass.cute.runtime import from_dlpack\n", - "\n", - "\n", - "@cute.jit\n", - "def print_tensor_dlpack(src: cute.Tensor):\n", - " print(src)\n", - " cute.print_tensor(src)" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "a = torch.randn(8, 5, dtype=torch_dtype(cutlass.Float32))\n", - "\n", - "print_tensor_dlpack(from_dlpack(a))" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "import numpy as np\n", - "\n", - "a = np.random.randn(8, 8).astype(np.float32)\n", - "\n", - "print_tensor_dlpack(from_dlpack(a))" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Tensor Evaluation Methods\n", - "\n", - "Tensors support two primary methods of evaluation:\n", - "\n", - "### 1. Full Evaluation\n", - "When applying the tensor evaluation with a complete coordinate c, it computes the offset, applies it to the engine, \n", - "and dereferences it to return the stored value. This is the straightforward case where you want to access \n", - "a specific element of the tensor.\n", - "\n", - "### 2. Partial Evaluation (Slicing)\n", - "When evaluating with an incomplete coordinate c = c' ⊕ c* (where c* represents the unspecified portion), \n", - "the result is a new tensor which is a slice of the original tensor with its engine offset to account for \n", - "the coordinates that were provided. This operation can be expressed as:\n", - "\n", - "```\n", - "T(c) = (E ∘ L)(c) = (E + L(c')) ∘ L(c*) = T'(c*)\n", - "```\n", - "\n", - "Slicing effectively reduces the dimensionality of the tensor, creating a sub-tensor that can be \n", - "further evaluated or manipulated." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def tensor_access_item(a: cute.Tensor):\n", - " # access data using linear index\n", - " cute.printf(\n", - " \"a[2] = {} (equivalent to a[{}])\",\n", - " a[2],\n", - " cute.make_identity_tensor(a.layout.shape)[2],\n", - " )\n", - " cute.printf(\n", - " \"a[9] = {} (equivalent to a[{}])\",\n", - " a[9],\n", - " cute.make_identity_tensor(a.layout.shape)[9],\n", - " )\n", - "\n", - " # access data using n-d coordinates, following two are equivalent\n", - " cute.printf(\"a[2,0] = {}\", a[2, 0])\n", - " cute.printf(\"a[2,4] = {}\", a[2, 4])\n", - " cute.printf(\"a[(2,4)] = {}\", a[2, 4])\n", - "\n", - " # assign value to tensor@(2,4)\n", - " a[2, 3] = 100.0\n", - " a[2, 4] = 101.0\n", - " cute.printf(\"a[2,3] = {}\", a[2, 3])\n", - " cute.printf(\"a[(2,4)] = {}\", a[(2, 4)])\n", - "\n", - "\n", - "# Create a tensor with sequential data using torch\n", - "data = torch.arange(0, 8 * 5, dtype=torch.float32).reshape(8, 5)\n", - "tensor_access_item(from_dlpack(data))\n", - "\n", - "print(data)" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "### Tensor as memory view\n", - "\n", - "In CUDA programming, different memory spaces have different characteristics in terms of access speed, scope, and lifetime:\n", - "\n", - "- **generic**: Default memory space that can refer to any other memory space.\n", - "- **global memory (gmem)**: Accessible by all threads across all blocks, but has higher latency.\n", - "- **shared memory (smem)**: Accessible by all threads within a block, with much lower latency than global memory.\n", - "- **register memory (rmem)**: Thread-private memory with the lowest latency, but limited capacity.\n", - "- **tensor memory (tmem)**: Specialized memory introduced in NVIDIA Blackwell architecture for tensor operations.\n", - "\n", - "When creating tensors in CuTe, you can specify the memory space to optimize performance based on your access patterns.\n", - "\n", - "For more information on CUDA memory spaces, see the [CUDA Programming Guide](https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#memory-hierarchy).\n" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Coordinate Tensors\n", - "\n", - "### Definition and Properties\n", - "\n", - "A coordinate tensor $T: Z^n → Z^m$ is a mathematical structure that establishes a mapping between coordinate spaces. Unlike standard tensors that map coordinates to scalar values, coordinate tensors map coordinates to other coordinates, forming a fundamental building block for tensor operations and transformations.\n", - "\n", - "### Examples\n", - "\n", - "Consider a `(4,4)` coordinate tensor:\n", - "\n", - "**Row-Major Layout (C-style):**\n", - "\\begin{bmatrix} \n", - "(0,0) & (0,1) & (0,2) & (0,3) \\\\\n", - "(1,0) & (1,1) & (1,2) & (1,3) \\\\\n", - "(2,0) & (2,1) & (2,2) & (2,3) \\\\\n", - "(3,0) & (3,1) & (3,2) & (3,3)\n", - "\\end{bmatrix}\n", - "\n", - "**Column-Major Layout (Fortran-style):**\n", - "\\begin{bmatrix}\n", - "(0,0) & (1,0) & (2,0) & (3,0) \\\\\n", - "(0,1) & (1,1) & (2,1) & (3,1) \\\\\n", - "(0,2) & (1,2) & (2,2) & (3,2) \\\\\n", - "(0,3) & (1,3) & (2,3) & (3,3)\n", - "\\end{bmatrix}\n", - "\n", - "### Identity Tensor\n", - "\n", - "An identity tensor $I$ is a special case of a coordinate tensor that implements the identity mapping function:\n", - "\n", - "**Definition:**\n", - "For a given shape $S = (s_1, s_2, ..., s_n)$, the identity tensor $I$ satisfies: $I(c) = c, \\forall c \\in \\prod_{i=1}^n [0, s_i)$\n", - "\n", - "**Properties:**\n", - "1. **Bijective Mapping**: The identity tensor establishes a one-to-one correspondence between coordinates.\n", - "2. **Layout Invariance**: The logical structure remains constant regardless of the underlying memory layout.\n", - "3. **Coordinate Preservation**: For any coordinate c, I(c) = c.\n", - "\n", - "\n", - "CuTe establishes an isomorphism between 1-D indices and N-D coordinates through lexicographical ordering. For a coordinate c = (c₁, c₂, ..., cₙ) in an identity tensor with shape S = (s₁, s₂, ..., sₙ):\n", - "\n", - "**Linear Index Formula:**\n", - "$\\text{idx} = c_1 + \\sum_{i=2}^{n} \\left(c_i \\prod_{j=1}^{i-1} s_j\\right)$\n", - "\n", - "**Example:**\n", - "```python\n", - "# Create an identity tensor from a given shape\n", - "coord_tensor = make_identity_tensor(layout.shape())\n", - "\n", - "# Access coordinate using linear index\n", - "coord = coord_tensor[linear_idx] # Returns the N-D coordinate\n", - "```\n", - "\n", - "This bidirectional mapping enables efficient conversion from linear indices to N-dimensional coordinates, facilitating tensor operations and memory access patterns." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def print_tensor_coord(a: cute.Tensor):\n", - " coord_tensor = cute.make_identity_tensor(a.layout.shape)\n", - " print(coord_tensor)\n", - " cute.print_tensor(coord_tensor)\n", - "\n", - "\n", - "a = torch.randn(8, 4, dtype=torch_dtype(cutlass.Float32))\n", - "print_tensor_coord(from_dlpack(a))" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": "Python 3 (ipykernel)", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.12.5" - }, - "widgets": { - "application/vnd.jupyter.widget-state+json": { - "state": {}, - "version_major": 2, - "version_minor": 0 - } - } - }, - "nbformat": 4, - "nbformat_minor": 4 -} diff --git a/examples/python/CuTeDSL/cute/notebooks/tensorssa.ipynb b/examples/python/CuTeDSL/cute/notebooks/tensorssa.ipynb deleted file mode 100644 index 62804a10db..0000000000 --- a/examples/python/CuTeDSL/cute/notebooks/tensorssa.ipynb +++ /dev/null @@ -1,495 +0,0 @@ -{ - "cells": [ - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "import cutlass\n", - "import cutlass.cute as cute\n", - "from cutlass.cute.runtime import from_dlpack\n", - "\n", - "import numpy as np" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "# Introduction to the TensorSSA in CuTe DSL\n", - "\n", - "This tutorial introduces what is the `TensorSSA` and why we need it. We also give some examples to show how to use `TensorSSA`.\n", - "\n", - "## What is TensorSSA\n", - "\n", - "`TensorSSA` is a Python class that represents a tensor value in Static Single Assignment (SSA) form within the CuTe DSL. You can think of it as a tensor residing in a (simulated) register.\n", - "\n", - "## Why TensorSSA\n", - "\n", - "`TensorSSA` encapsulates the underlying MLIR tensor value into an object that's easier to manipulate in Python. By overloading numerous Python operators (like `+`, `-`, `*`, `/`, `[]`, etc.), it allows users to express tensor computations (primarily element-wise operations and reductions) in a more Pythonic way. These element-wise operations are then translated into optimized vectorization instructions.\n", - "\n", - "It's part of the CuTe DSL, serving as a bridge between the user-described computational logic and the lower-level MLIR IR, particularly for representing and manipulating register-level data.\n", - "\n", - "## When to use TensorSSA\n", - "\n", - "`TensorSSA` is primarily used in the following scenarios:\n", - "\n", - "### Load from memory and store to memory" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def load_and_store(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):\n", - " \"\"\"\n", - " Load data from memory and store the result to memory.\n", - "\n", - " :param res: The destination tensor to store the result.\n", - " :param a: The source tensor to be loaded.\n", - " :param b: The source tensor to be loaded.\n", - " \"\"\"\n", - " a_vec = a.load()\n", - " print(f\"a_vec: {a_vec}\") # prints `a_vec: vector<12xf32> o (3, 4)`\n", - " b_vec = b.load()\n", - " print(f\"b_vec: {b_vec}\") # prints `b_vec: vector<12xf32> o (3, 4)`\n", - " res.store(a_vec + b_vec)\n", - " cute.print_tensor(res)\n", - "\n", - "\n", - "a = np.ones(12).reshape((3, 4)).astype(np.float32)\n", - "b = np.ones(12).reshape((3, 4)).astype(np.float32)\n", - "c = np.zeros(12).reshape((3, 4)).astype(np.float32)\n", - "load_and_store(from_dlpack(c), from_dlpack(a), from_dlpack(b))" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "### Register-Level Tensor Operations\n", - "\n", - "When writing kernel logic, various computations, transformations, slicing, etc., are performed on data loaded into registers." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def apply_slice(src: cute.Tensor, dst: cute.Tensor, indices: cutlass.Constexpr):\n", - " \"\"\"\n", - " Apply slice operation on the src tensor and store the result to the dst tensor.\n", - "\n", - " :param src: The source tensor to be sliced.\n", - " :param dst: The destination tensor to store the result.\n", - " :param indices: The indices to slice the source tensor.\n", - " \"\"\"\n", - " src_vec = src.load()\n", - " dst_vec = src_vec[indices]\n", - " print(f\"{src_vec} -> {dst_vec}\")\n", - " if cutlass.const_expr(isinstance(dst_vec, cute.TensorSSA)):\n", - " dst.store(dst_vec)\n", - " cute.print_tensor(dst)\n", - " else:\n", - " dst[0] = dst_vec\n", - " cute.print_tensor(dst)\n", - "\n", - "\n", - "def slice_1():\n", - " src_shape = (4, 2, 3)\n", - " dst_shape = (4, 3)\n", - " indices = (None, 1, None)\n", - "\n", - " \"\"\"\n", - " a:\n", - " [[[ 0. 1. 2.]\n", - " [ 3. 4. 5.]]\n", - "\n", - " [[ 6. 7. 8.]\n", - " [ 9. 10. 11.]]\n", - "\n", - " [[12. 13. 14.]\n", - " [15. 16. 17.]]\n", - "\n", - " [[18. 19. 20.]\n", - " [21. 22. 23.]]]\n", - " \"\"\"\n", - " a = np.arange(np.prod(src_shape)).reshape(*src_shape).astype(np.float32)\n", - " dst = np.random.randn(*dst_shape).astype(np.float32)\n", - " apply_slice(from_dlpack(a), from_dlpack(dst), indices)\n", - "\n", - "\n", - "slice_1()" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "def slice_2():\n", - " src_shape = (4, 2, 3)\n", - " dst_shape = (1,)\n", - " indices = 10\n", - " a = np.arange(np.prod(src_shape)).reshape(*src_shape).astype(np.float32)\n", - " dst = np.random.randn(*dst_shape).astype(np.float32)\n", - " apply_slice(from_dlpack(a), from_dlpack(dst), indices)\n", - "\n", - "\n", - "slice_2()" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Arithmetic Operations\n", - "\n", - "As we mentioned earlier, there're many tensor operations whose operands are `TensorSSA`. And they are all element-wise operations. We give some examples below.\n", - "\n", - "### Binary Operations\n", - "\n", - "For binary operations, the LHS operand is `TensorSSA` and the RHS operand can be either `TensorSSA` or `Numeric`. When the RHS is `Numeric`, it will be broadcast to a `TensorSSA`." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def binary_op_1(a: cute.Tensor, b: cute.Tensor):\n", - " a_vec = a.load()\n", - " b_vec = b.load()\n", - "\n", - " add_res = a_vec + b_vec\n", - " cute.print_tensor(add_res) # prints [3.000000, 3.000000, 3.000000]\n", - "\n", - " sub_res = a_vec - b_vec\n", - " cute.print_tensor(sub_res) # prints [-1.000000, -1.000000, -1.000000]\n", - "\n", - " mul_res = a_vec * b_vec\n", - " cute.print_tensor(mul_res) # prints [2.000000, 2.000000, 2.000000]\n", - "\n", - " div_res = a_vec / b_vec\n", - " cute.print_tensor(div_res) # prints [0.500000, 0.500000, 0.500000]\n", - "\n", - " floor_div_res = a_vec // b_vec\n", - " cute.print_tensor(floor_div_res) # prints [0.000000, 0.000000, 0.000000]\n", - "\n", - " mod_res = a_vec % b_vec\n", - " cute.print_tensor(mod_res) # prints [1.000000, 1.000000, 1.000000]\n", - "\n", - "\n", - "a = np.empty((3,), dtype=np.float32)\n", - "a.fill(1.0)\n", - "b = np.empty((3,), dtype=np.float32)\n", - "b.fill(2.0)\n", - "binary_op_1(from_dlpack(a), from_dlpack(b))" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def binary_op_2(a: cute.Tensor, c: cutlass.Constexpr):\n", - " a_vec = a.load()\n", - "\n", - " add_res = a_vec + c\n", - " cute.print_tensor(add_res) # prints [3.000000, 3.000000, 3.000000]\n", - "\n", - " sub_res = a_vec - c\n", - " cute.print_tensor(sub_res) # prints [-1.000000, -1.000000, -1.000000]\n", - "\n", - " mul_res = a_vec * c\n", - " cute.print_tensor(mul_res) # prints [2.000000, 2.000000, 2.000000]\n", - "\n", - " div_res = a_vec / c\n", - " cute.print_tensor(div_res) # prints [0.500000, 0.500000, 0.500000]\n", - "\n", - " floor_div_res = a_vec // c\n", - " cute.print_tensor(floor_div_res) # prints [0.000000, 0.000000, 0.000000]\n", - "\n", - " mod_res = a_vec % c\n", - " cute.print_tensor(mod_res) # prints [1.000000, 1.000000, 1.000000]\n", - "\n", - "\n", - "a = np.empty((3,), dtype=np.float32)\n", - "a.fill(1.0)\n", - "c = 2.0\n", - "binary_op_2(from_dlpack(a), c)" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def binary_op_3(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):\n", - " a_vec = a.load()\n", - " b_vec = b.load()\n", - "\n", - " gt_res = a_vec > b_vec\n", - " res.store(gt_res)\n", - "\n", - " \"\"\"\n", - " ge_res = a_ >= b_ # [False, True, False]\n", - " lt_res = a_ < b_ # [True, False, True]\n", - " le_res = a_ <= b_ # [True, False, True]\n", - " eq_res = a_ == b_ # [False, False, False]\n", - " \"\"\"\n", - "\n", - "\n", - "a = np.array([1, 2, 3], dtype=np.float32)\n", - "b = np.array([2, 1, 4], dtype=np.float32)\n", - "res = np.empty((3,), dtype=np.bool_)\n", - "binary_op_3(from_dlpack(res), from_dlpack(a), from_dlpack(b))\n", - "print(res) # prints [False, True, False]" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def binary_op_4(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):\n", - " a_vec = a.load()\n", - " b_vec = b.load()\n", - "\n", - " xor_res = a_vec ^ b_vec\n", - " res.store(xor_res)\n", - "\n", - " # or_res = a_vec | b_vec\n", - " # res.store(or_res) # prints [3, 2, 7]\n", - "\n", - " # and_res = a_vec & b_vec\n", - " # res.store(and_res) # prints [0, 2, 0]\n", - "\n", - "\n", - "a = np.array([1, 2, 3], dtype=np.int32)\n", - "b = np.array([2, 2, 4], dtype=np.int32)\n", - "res = np.empty((3,), dtype=np.int32)\n", - "binary_op_4(from_dlpack(res), from_dlpack(a), from_dlpack(b))\n", - "print(res) # prints [3, 0, 7]" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "#### Unary Operations" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def unary_op_1(res: cute.Tensor, a: cute.Tensor):\n", - " a_vec = a.load()\n", - "\n", - " sqrt_res = cute.math.sqrt(a_vec)\n", - " cute.print_tensor(sqrt_res) # prints [2.000000, 2.000000, 2.000000]\n", - "\n", - " sin_res = cute.math.sin(a_vec)\n", - " res.store(sin_res)\n", - " cute.print_tensor(sin_res) # prints [-0.756802, -0.756802, -0.756802]\n", - "\n", - " exp2_res = cute.math.exp2(a_vec)\n", - " cute.print_tensor(exp2_res) # prints [16.000000, 16.000000, 16.000000]\n", - "\n", - "\n", - "a = np.array([4.0, 4.0, 4.0], dtype=np.float32)\n", - "res = np.empty((3,), dtype=np.float32)\n", - "unary_op_1(from_dlpack(res), from_dlpack(a))" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "#### Reduction Operation\n", - "\n", - "The `TensorSSA`'s `reduce` method applies a specified reduction operation (`ReductionOp.ADD`, \n", - "`ReductionOp.MUL`, `ReductionOp.MAX`, `ReductionOp.MIN`) starting with an initial value, and \n", - "performs this reduction along the dimensions specified by the `reduction_profile`. The result \n", - "is typically a new `TensorSSA` with reduced dimensions or a scalar value if it reduces across \n", - "all axes." - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "@cute.jit\n", - "def reduction_op(a: cute.Tensor):\n", - " \"\"\"\n", - " Apply reduction operation on the src tensor.\n", - "\n", - " :param src: The source tensor to be reduced.\n", - " \"\"\"\n", - " a_vec = a.load()\n", - " red_res = a_vec.reduce(cute.ReductionOp.ADD, 0.0, reduction_profile=0)\n", - " cute.printf(red_res) # prints 21.000000\n", - "\n", - " red_res = a_vec.reduce(cute.ReductionOp.ADD, 0.0, reduction_profile=(None, 1))\n", - " cute.print_tensor(red_res) # prints [6.000000, 15.000000]\n", - "\n", - " red_res = a_vec.reduce(cute.ReductionOp.ADD, 1.0, reduction_profile=(1, None))\n", - " cute.print_tensor(red_res) # prints [6.000000, 8.000000, 10.000000]\n", - "\n", - "\n", - "a = np.array([[1, 2, 3], [4, 5, 6]], dtype=np.float32)\n", - "reduction_op(from_dlpack(a))" - ] - }, - { - "cell_type": "markdown", - "metadata": {}, - "source": [ - "## Broadcast\n", - "\n", - "`TensorSSA` supports broadcasting operations following NumPy's broadcasting rules. Broadcasting \n", - "allows you to perform operations on arrays of different shapes when certain conditions are met. \n", - "The key rules are:\n", - "\n", - "1. Source shape is padded with 1's to match the rank of target shape\n", - "2. The size in each mode of source shape must either be 1 or equal to target shape\n", - "3. After broadcasting, all modes should match target shape\n", - "\n", - "Let's look at some examples of broadcasting in action:" - ] - }, - { - "cell_type": "code", - "execution_count": null, - "metadata": {}, - "outputs": [], - "source": [ - "import cutlass\n", - "import cutlass.cute as cute\n", - "\n", - "\n", - "@cute.jit\n", - "def broadcast_examples():\n", - " a = cute.make_rmem_tensor((1, 3), dtype=cutlass.Float32)\n", - " a[0] = 0.0\n", - " a[1] = 1.0\n", - " a[2] = 2.0\n", - " a_val = a.load()\n", - " cute.print_tensor(a_val.broadcast_to((4, 3)))\n", - " # tensor(raw_ptr(0x00007ffe26625740: f32, rmem, align<32>) o (4,3):(1,4), data=\n", - " # [[ 0.000000, 1.000000, 2.000000, ],\n", - " # [ 0.000000, 1.000000, 2.000000, ],\n", - " # [ 0.000000, 1.000000, 2.000000, ],\n", - " # [ 0.000000, 1.000000, 2.000000, ]])\n", - "\n", - " c = cute.make_rmem_tensor((4, 1), dtype=cutlass.Float32)\n", - " c[0] = 0.0\n", - " c[1] = 1.0\n", - " c[2] = 2.0\n", - " c[3] = 3.0\n", - " cute.print_tensor(a.load() + c.load())\n", - " # tensor(raw_ptr(0x00007ffe26625780: f32, rmem, align<32>) o (4,3):(1,4), data=\n", - " # [[ 0.000000, 1.000000, 2.000000, ],\n", - " # [ 1.000000, 2.000000, 3.000000, ],\n", - " # [ 2.000000, 3.000000, 4.000000, ],\n", - " # [ 3.000000, 4.000000, 5.000000, ]])\n", - "\n", - "\n", - "broadcast_examples()" - ] - }, - { - "cell_type": "markdown", - "metadata": { - "vscode": { - "languageId": "raw" - } - }, - "source": [ - "The examples above demonstrate two key broadcasting scenarios:\n", - "\n", - "1. **Row Vector Broadcasting**: In the first example, we create a row vector `a` with shape \n", - " (1, 3) containing values [0.0, 1.0, 2.0]. When we broadcast it to shape (4, 3), the values \n", - " are repeated across the first dimension, resulting in:\n", - " ```\n", - " [[0.0, 1.0, 2.0],\n", - " [0.0, 1.0, 2.0],\n", - " [0.0, 1.0, 2.0],\n", - " [0.0, 1.0, 2.0]]\n", - " ```\n", - " This demonstrates how a row vector can be broadcast to create multiple identical rows.\n", - "\n", - "2. **Column Vector and Row Vector Addition**: In the second example, we have:\n", - " - A row vector `a` with shape (1, 3) containing [0.0, 1.0, 2.0]\n", - " - A column vector `c` with shape (4, 1) containing [0.0, 1.0, 2.0, 3.0]\n", - " \n", - " When we add these together, both vectors are broadcast to shape (4, 3):\n", - " - The row vector is broadcast vertically (4 times)\n", - " - The column vector is broadcast horizontally (3 times)\n", - " \n", - " The result is:\n", - " ```\n", - " [[0.0 + 0.0, 1.0 + 0.0, 2.0 + 0.0],\n", - " [0.0 + 1.0, 1.0 + 1.0, 2.0 + 1.0],\n", - " [0.0 + 2.0, 1.0 + 2.0, 2.0 + 2.0],\n", - " [0.0 + 3.0, 1.0 + 3.0, 2.0 + 3.0]]\n", - " ```\n", - " =\n", - " ```\n", - " [[0.0, 1.0, 2.0],\n", - " [1.0, 2.0, 3.0],\n", - " [2.0, 3.0, 4.0],\n", - " [3.0, 4.0, 5.0]]\n", - " ```\n", - "\n", - "This demonstrates how `TensorSSA` can automatically handle broadcasting of both row and column \n", - "vectors in arithmetic operations, following the broadcasting rules where each dimension must \n", - "either be 1 or match the target size. The broadcasting is handled implicitly during operations, \n", - "making it easy to work with tensors of different shapes.\n" - ] - } - ], - "metadata": { - "kernelspec": { - "display_name": ".venv3_12", - "language": "python", - "name": "python3" - }, - "language_info": { - "codemirror_mode": { - "name": "ipython", - "version": 3 - }, - "file_extension": ".py", - "mimetype": "text/x-python", - "name": "python", - "nbconvert_exporter": "python", - "pygments_lexer": "ipython3", - "version": "3.12.11" - } - }, - "nbformat": 4, - "nbformat_minor": 4 -} diff --git a/examples/python/CuTeDSL/cute/rubin/kernel/dense_gemm/dense_gemm_persistent_mixed_clusters.py b/examples/python/CuTeDSL/cute/rubin/kernel/dense_gemm/dense_gemm_persistent_mixed_clusters.py index b276e18109..9da809ffe5 100644 --- a/examples/python/CuTeDSL/cute/rubin/kernel/dense_gemm/dense_gemm_persistent_mixed_clusters.py +++ b/examples/python/CuTeDSL/cute/rubin/kernel/dense_gemm/dense_gemm_persistent_mixed_clusters.py @@ -404,8 +404,14 @@ def _compute_grid( # otherwise when the division is not exact, the extra partial wave of # preferred clusters may force the hardware to schedule one additional wave, # causing significant performance regression. - max_preferred_cluster_count = ( - max_ctas_for_fallback_cluster // preferred_cluster_size_mn + # It is possible to end up with a small problem size where + # max_ctas_for_fallback_cluster < preferred_cluster_size_mn, which leads to + # max_preferred_cluster_count being zero; in this case, we + # cannot even fill a single wave, and so we are not concerned with performance + # implications rather we choose to still be able to execute properly instead of + # erroring out. + max_preferred_cluster_count = cutlass.max( + 1, max_ctas_for_fallback_cluster // preferred_cluster_size_mn ) preferred_grid = ( preferred_grid[0], @@ -506,19 +512,6 @@ def can_implement( f"integer multiple of fallback cluster shape {fallback_cluster_shape_mn}" ) - # Check that the problem is at least as large as the preferred cluster tile. - # The mixed clusters kernel computes max_preferred_cluster_count as: - # max_ctas_for_fallback_cluster // preferred_cluster_size_mn - # If the problem is smaller than one preferred cluster tile, this count - # becomes zero, resulting in an invalid grid shape. - m, n, k, l = mnkl - preferred_tile_m = mma_tiler[0] * preferred_cluster_shape_mn[0] - preferred_tile_n = mma_tiler[1] * preferred_cluster_shape_mn[1] - if m < preferred_tile_m or n < preferred_tile_n: - raise testing.CantImplementError( - f"Problem size ({m}, {n}) is smaller than the preferred cluster tile " - f"({preferred_tile_m}, {preferred_tile_n})" - ) except testing.CantImplementError as e: print(f"[DSL ERROR] CantImplementError: {e}") return False diff --git a/examples/python/CuTeDSL/cute_ext/blackwell/dense_gemm/dense_gemm_alpha_beta_persistent.py b/examples/python/CuTeDSL/cute_ext/blackwell/dense_gemm/dense_gemm_alpha_beta_persistent.py index b91e217e6d..37eb08edee 100644 --- a/examples/python/CuTeDSL/cute_ext/blackwell/dense_gemm/dense_gemm_alpha_beta_persistent.py +++ b/examples/python/CuTeDSL/cute_ext/blackwell/dense_gemm/dense_gemm_alpha_beta_persistent.py @@ -695,7 +695,7 @@ def kernel( acc_epi_div_tiled = cute.flat_divide(accumulators_sliced, epi_tile) subtile_cnt = cute.size(acc_epi_div_tiled.shape, mode=[3]) - for mn in range(subtile_cnt): + for mn in cutlass.range(subtile_cnt, unroll_full=True): # TMEM -> RMEM cute_ext.partition_and_copy( tiled_copy_t2r.get_slice(tid_x), diff --git a/examples/python/CuTeDSL/cute_ext/rubin/dense_blockscaled_gemm_persistent.py b/examples/python/CuTeDSL/cute_ext/rubin/dense_blockscaled_gemm_persistent.py index b5c1227f4a..7dc9ef41a9 100755 --- a/examples/python/CuTeDSL/cute_ext/rubin/dense_blockscaled_gemm_persistent.py +++ b/examples/python/CuTeDSL/cute_ext/rubin/dense_blockscaled_gemm_persistent.py @@ -33,13 +33,7 @@ from typing import Optional, Type, Tuple, Literal, Callable import cutlass -from cutlass import core -from cutlass import utils -import cutlass.cute as cute -from cutlass import ( - cute as cute, - utils as utils, -) +from cutlass import core, cute, utils from cutlass.cute.nvgpu import tcgen05 from cutlass.cute.nvgpu.tcgen05.mma import CollectorOp @@ -59,9 +53,6 @@ # b-reuse where ``MMA_M == 2``. from cutlass.utils.gemm.sm100 import transform_partitioned_tensor_layout -# Persistent tile scheduler lives in the local ``helpers`` package alongside -# the cute_ext examples (rooted at ``CuTeDSL/``); the cute_ext// files -# add ``CuTeDSL/`` to ``sys.path`` so the import resolves. if __name__ == "__main__": current_dir = os.path.dirname(os.path.abspath(__file__)) sys.path.insert(0, os.path.join(current_dir, "..")) @@ -88,9 +79,10 @@ (it has FP4-specific entries and falls back to the sm_100 table otherwise). * ``arch="sm_107"`` selects Rubin's larger SMEM budget. * The MMA tiler K is supplied separately from the MMA instruction K via - ``--mma_tiler`` / ``--mma_inst_shape``. For Rubin FP4 the atom's K is 128 - (vs sm_100's 64), so the typical configuration is mma_tiler_k=256 and - mma_inst_shape_k=128, keeping the per-stage K tile at 256. + ``--mma_tiler`` / ``--mma_inst_shape``. This kernel uses instruction/tiler + K=64/128 for FP8xFP8 and mixed FP8/FP4 configurations. Pure FP4 x FP4 uses + instruction K=128 and preserves the existing positive-multiple-of-128 tiler + envelope; the reference configuration uses tiler K=256. The LIR-level operation types (``SM100_MMA_SCALED_*``, ``SM90_TMA_LOAD``, ``cute_ext.dot_block_scaled``, ...) are functionally a superset of what sm_107 @@ -98,9 +90,11 @@ Blackwell LIR version once the dust settles -- duplication is intentional for now. -Initial scope : +Supported scope: -- A/B dtype: ``Float4E2M1FN`` only. +- A/B dtype: all ``Float8E4M3FN`` / ``Float8E5M2`` pairings, + ``Float4E2M1FN`` for both operands, and mixed FP8/FP4 in either + operand order. - B-reuse: opt-in via ``mma_tiler_m // mma_inst_shape_m == 2``. The MMA warp then issues a bkeep / breuse pair per K-block (FILL / LASTUSE collector ops on B), sharing one B+SFB load between the two M-halves. @@ -257,16 +251,102 @@ def can_implement( f"D major axis {c_major!r} not supported; expected one of ('n', 'm')." ) - # ---- Rubin initial-scope restrictions ------------------------------- - if a_dtype is not cutlass.Float4E2M1FN: - return ( - f"SM107 LIR initial scope: only a_dtype=Float4E2M1FN is " - f"supported; got {a_dtype}." - ) - if b_dtype is not cutlass.Float4E2M1FN: + supported_dtype_configs = { + ( + cutlass.Float8E4M3FN, + cutlass.Float8E4M3FN, + cutlass.Float8E8M0FNU, + 32, + ), + ( + cutlass.Float8E4M3FN, + cutlass.Float8E5M2, + cutlass.Float8E8M0FNU, + 32, + ), + ( + cutlass.Float8E5M2, + cutlass.Float8E4M3FN, + cutlass.Float8E8M0FNU, + 32, + ), + ( + cutlass.Float8E5M2, + cutlass.Float8E5M2, + cutlass.Float8E8M0FNU, + 32, + ), + ( + cutlass.Float4E2M1FN, + cutlass.Float4E2M1FN, + cutlass.Float8E4M3FN, + 16, + ), + ( + cutlass.Float4E2M1FN, + cutlass.Float4E2M1FN, + cutlass.Float8E4M3FN, + 32, + ), + ( + cutlass.Float4E2M1FN, + cutlass.Float4E2M1FN, + cutlass.Float8E8M0FNU, + 16, + ), + ( + cutlass.Float4E2M1FN, + cutlass.Float4E2M1FN, + cutlass.Float8E8M0FNU, + 32, + ), + ( + cutlass.Float4E2M1FN, + cutlass.Float4E2M1FN, + cutlass.FloatNV8E5M3FNU, + 16, + ), + ( + cutlass.Float4E2M1FN, + cutlass.Float4E2M1FN, + cutlass.FloatNV8E5M3FNU, + 32, + ), + ( + cutlass.Float8E4M3FN, + cutlass.Float4E2M1FN, + cutlass.Float8E8M0FNU, + 32, + ), + ( + cutlass.Float8E5M2, + cutlass.Float4E2M1FN, + cutlass.Float8E8M0FNU, + 32, + ), + ( + cutlass.Float4E2M1FN, + cutlass.Float8E4M3FN, + cutlass.Float8E8M0FNU, + 32, + ), + ( + cutlass.Float4E2M1FN, + cutlass.Float8E5M2, + cutlass.Float8E8M0FNU, + 32, + ), + } + if (a_dtype, b_dtype, sf_dtype, sf_vec_size) not in supported_dtype_configs: return ( - f"SM107 LIR initial scope: only b_dtype=Float4E2M1FN is " - f"supported; got {b_dtype}." + f"Unsupported (a_dtype, b_dtype, sf_dtype, sf_vec_size) " + f"combination: ({a_dtype}, {b_dtype}, {sf_dtype}, {sf_vec_size}); " + f"expected FP8 with A/B in " + f"{{Float8E4M3FN, Float8E5M2}}, Float8E8M0FNU scale factors, " + f"and sf_vec_size=32; mixed FP8/FP4 with the same scale-factor " + f"configuration; or FP4 with sf_dtype in " + f"{{Float8E8M0FNU, Float8E4M3FN, FloatNV8E5M3FNU}} and " + f"sf_vec_size in {{16, 32}}." ) if mma_inst_shape[0] not in (128, 256): return ( @@ -306,14 +386,23 @@ def can_implement( ) # --------------------------------------------------------------------- - # FP4 atom's K is fixed at 128. - if mma_inst_shape[2] != 128: - return f"FP4 MMA instruction K must be 128; got {mma_inst_shape[2]}." - # MMA tiler K must be a positive multiple of the MMA instruction K. - if mma_tiler[2] <= 0 or mma_tiler[2] % mma_inst_shape[2] != 0: + is_fp4 = a_dtype is cutlass.Float4E2M1FN and b_dtype is cutlass.Float4E2M1FN + if is_fp4 and mma_inst_shape[2] != 128: + return f"FP4 requires mma_inst_shape_k=128; got {mma_inst_shape[2]}." + # Preserve the kernel's existing support for any positive number of + # FP4 instruction-K blocks per load stage. This kernel keeps pure FP4 + # at instruction K=128; the reference uses two blocks (tiler K=256). + if is_fp4 and (mma_tiler[2] <= 0 or mma_tiler[2] % mma_inst_shape[2] != 0): + return ( + f"FP4 mma_tiler_k={mma_tiler[2]} must be a positive multiple " + f"of mma_inst_shape_k={mma_inst_shape[2]}." + ) + if not is_fp4 and (mma_inst_shape[2] != 64 or mma_tiler[2] != 128): return ( - f"mma_tiler_k={mma_tiler[2]} must be a positive multiple of " - f"mma_inst_shape_k={mma_inst_shape[2]}." + f"This SM107 block-scaled kernel requires mma_inst_shape_k=64 " + f"and mma_tiler_k=128 for MXF8F6F4 configurations; got " + f"mma_inst_shape_k={mma_inst_shape[2]} and " + f"mma_tiler_k={mma_tiler[2]}." ) if mma_inst_shape[1] not in (64, 128, 192, 256): @@ -321,25 +410,12 @@ def can_implement( f"MMA instruction N {mma_inst_shape[1]} not supported; " f"expected one of (64, 128, 192, 256)." ) - # FP4 SF dtype / vec size combos supported by SM107MmaMXF4NVF4Op. - # TODO: add E5M3 support - if (sf_dtype, sf_vec_size) not in ( - (cutlass.Float8E4M3FN, 16), - (cutlass.Float8E4M3FN, 32), - (cutlass.Float8E8M0FNU, 16), - (cutlass.Float8E8M0FNU, 32), - ): - return ( - f"Unsupported (sf_dtype, sf_vec_size) combination: " - f"({sf_dtype}, {sf_vec_size}) for FP4; supported combos: " - f"(Float8E4M3FN/Float8E8M0FNU, 16/32)." - ) - # FP4 atoms require K-major A/B and N-major D. - if not (a_major == b_major == "k" and c_major == "n"): - return ( - f"FP4 block-scaled requires a_major='k', b_major='k', c_major='n'; " - f"got a_major={a_major!r}, b_major={b_major!r}, c_major={c_major!r}." - ) + # Every FP4 operand must be K-major. FP8 operands support both major + # modes admitted at the start of this predicate. + if a_dtype is cutlass.Float4E2M1FN and a_major != "k": + return f"FP4 operand A requires a_major='k'; got a_major={a_major!r}." + if b_dtype is cutlass.Float4E2M1FN and b_major != "k": + return f"FP4 operand B requires b_major='k'; got b_major={b_major!r}." if c_dtype not in ( cutlass.Float16, cutlass.BFloat16, @@ -858,7 +934,6 @@ def epilogue_tma_store( # - All warps advance to the next pipeline stage tma_store_pipeline.release_advance() - tma_store_pipeline.tail() return tma_store_pipeline @cute.experimental.jit @@ -1123,8 +1198,7 @@ def kernel( cluster_layout_vmnk=cluster_layout_vmnk, ) - # UMMA -> tcgen05.ld. For 2-CTA both peers' epilogue warpgroups call - # ``consumer_release``, so the arrive count doubles. + # UMMA -> tcgen05.ld. acc_pipe = cute_ext.UMMAtoAsyncPipeline.create( num_stages=self.num_acc_stages, mma_operation_type=mma_operation_type, @@ -1490,6 +1564,10 @@ def kernel( tile_sched.advance_to_next_work() work_tile = tile_sched.get_current_work() + # Stage reuse across CTA tiles is protected by acquire_sync(), so + # the pipeline is drained once here instead of per CTA tile. + tma_store_pipe.tail() + # Cluster sync before exit. No-op for pure 1-CTA (1,1) launches. if cutlass.const_expr(cute.size(self.cluster_shape_mn) > 1): cute.arch.cluster_arrive() @@ -1755,8 +1833,8 @@ def run( output_tensor_name="C", ) - # Initial-scope gate -- delegate to the kernel's predicate so this stays in - # sync with what the kernel actually accepts. + # Delegate to the kernel's predicate so this stays in sync with what the + # kernel actually accepts. reason = Sm107BlockScaledDenseGemmKernel.can_implement( mnkl, mma_inst_shape, @@ -1890,7 +1968,8 @@ def generate_inputs(): type=cli.comma_separated_ints_of(3), default=(128, 128, 128), help="MMA instruction shape (M,N,K) comma-separated. mma_inst_m must " - "be 128 (1-CTA) or 256 (2-CTA); FP4 fixes mma_inst_k == 128.", + "be 128 (1-CTA) or 256 (2-CTA); mma_inst_k is 64 for FP8 and " + "mixed FP8/FP4, and 128 for pure FP4.", ) parser.add_argument( "--mma_tiler", @@ -1898,8 +1977,9 @@ def generate_inputs(): default=(128, 128, 256), help="MMA tiler shape (M,N,K) comma-separated. mma_tiler_m equals " "mma_inst_shape_m (no b-reuse) or 2*mma_inst_shape_m (b-reuse on); " - "mma_tiler_n must equal mma_inst_shape_n; mma_tiler_k must be a " - "positive multiple of mma_inst_shape_k.", + "mma_tiler_n must equal mma_inst_shape_n; mma_tiler_k must be " + "128 for FP8/mixed FP8-FP4 or a positive multiple of 128 for " + "pure FP4.", ) cli.add_cluster_shape_arg( parser, @@ -1913,9 +1993,9 @@ def generate_inputs(): parser.add_argument("--sf_dtype", type=cutlass.dtype, default=cutlass.Float8E8M0FNU) parser.add_argument("--sf_vec_size", type=int, default=32) cli.add_dtype_args(parser, c=cutlass.Float16) - # FP4 (initial scope) requires K-major A/B and N-major C; enforced by - # ``can_implement``. - cli.add_major_args(parser, a=["k"], b=["k"], c=["n"]) + # Every FP4 operand requires K-major; FP8 operands support the full matrix. + # The dtype-dependent constraints are enforced by ``can_implement``. + cli.add_major_args(parser) # Optional pipeline-depth caps; each only lowers the value picked by # ``_compute_stages`` (an override >= the computed depth has no effect). parser.add_argument("--num_load_stages_override", type=int, default=None) diff --git a/examples/python/CuTeDSL/experimental/task_scheduling/blackwell/tutorial/04_gemm_bf16_advanced_ts/02_fp16_bf16_gemm_3_cute_cluster.py b/examples/python/CuTeDSL/experimental/task_scheduling/blackwell/tutorial/04_gemm_bf16_advanced_ts/02_fp16_bf16_gemm_3_cute_cluster.py index c5ee62ba9f..34fba61969 100644 --- a/examples/python/CuTeDSL/experimental/task_scheduling/blackwell/tutorial/04_gemm_bf16_advanced_ts/02_fp16_bf16_gemm_3_cute_cluster.py +++ b/examples/python/CuTeDSL/experimental/task_scheduling/blackwell/tutorial/04_gemm_bf16_advanced_ts/02_fp16_bf16_gemm_3_cute_cluster.py @@ -91,17 +91,6 @@ from typing import Tuple, Optional, Type, Callable, Any from functools import partial, lru_cache from dataclasses import dataclass, field -from pathlib import Path -import sys - -_REPO_ROOT = Path(__file__).resolve().parents[10] -_repo_root_str = str(_REPO_ROOT) -if _repo_root_str not in sys.path: - sys.path.insert(0, _repo_root_str) -try: - import sitecustomize # noqa: F401 -except Exception: - pass import cutlass from cutlass import Numeric @@ -352,15 +341,15 @@ def __init__( self.cta_layout_size_v = cute.size(cta_layout_vmnk, mode=[0]) self.cta_layout_size_m = cute.size(cta_layout_vmnk, mode=[1]) self.cta_layout_size_n = cute.size(cta_layout_vmnk, mode=[2]) - self.act_num_pair_cols = ( + self.act_num_pair_cols = cutlass.Int32( act_num_pair_cols if act_num_pair_cols is not None else num_pair_cols ) - self.act_a_mcast_template = ( + self.act_a_mcast_template = cutlass.Int32( act_a_mcast_template if act_a_mcast_template is not None else _a_mcast_template ) - self.act_b_mcast_template = ( + self.act_b_mcast_template = cutlass.Int32( act_b_mcast_template if act_b_mcast_template is not None else _b_mcast_template @@ -391,6 +380,11 @@ def __init__( cute.recast_ptr(smem_ptr_b, self.b_smem_layout.inner), self.b_smem_layout.outer, ) + self.cta_rank_in_cluster = cutlass.Int32(0) + self.cta_in_cluster_coord_vmnk = (cutlass.Int32(0), cutlass.Int32(0), cutlass.Int32(0),cutlass.Int32(0)) + self.mma_v_coord = cutlass.Int32(0) + self.tma_mcast_mask_a = cutlass.Int16(0) + self.tma_mcast_mask_b = cutlass.Int16(0) def get_smem_requirements(self): return [self._alloc_a, self._alloc_b] @@ -621,6 +615,7 @@ def __init__( True, io_dtype.width, ) + self.cta_rank_in_cluster = cutlass.Int32(0) def get_tmem_requirements(self): return [self._alloc_acc] @@ -728,7 +723,7 @@ def load_subtile( tDtC_slice = tDtC_stage[(None, None, None, subtile_idx)] cute.copy(tiled_copy_t2r, tDtC_slice, tCrC) cute.arch.fence_view_async_tmem_load() - return tCrC.load() + return cutlass.Vector(tCrC.load()) @producer_work @cute.jit diff --git a/examples/python/CuTeDSL/notebooks/.gitignore b/examples/python/CuTeDSL/notebooks/.gitignore new file mode 100644 index 0000000000..31d2fd3502 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/.gitignore @@ -0,0 +1,4 @@ + +.ipynb_checkpoints/ +*.ncu-rep +playground.ipynb diff --git a/examples/python/CuTeDSL/notebooks/1_dsl_features/01_hello_world.ipynb b/examples/python/CuTeDSL/notebooks/1_dsl_features/01_hello_world.ipynb new file mode 100644 index 0000000000..9b3728b1e5 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/1_dsl_features/01_hello_world.ipynb @@ -0,0 +1,173 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Hello, world: your first CuTe DSL kernel\n", + "\n", + "The **CuTe DSL** lets you write a GPU kernel in Python and compile it straight to CUDA. Two\n", + "decorators do the work: **`@cute.kernel`** marks the *device* code that every GPU thread runs, and\n", + "**`@cute.jit`** marks the *host* function that launches it. You hand the kernel plain PyTorch CUDA\n", + "tensors and read the result straight back -- no C++, no separate build step.\n", + "\n", + "This first chapter is deliberately tiny: one kernel that adds two arrays and prints a greeting from\n", + "the GPU. Every later chapter builds on this shape." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** the two decorators -- `@cute.kernel` (device) and `@cute.jit` (host launcher); how\n", + "a thread finds its position with `cute.arch.thread_idx()` / `block_idx()` / `block_dim()`; how to\n", + "print from the device with `cute.printf`; and the cutlass <-> PyTorch round trip via\n", + "`cute.runtime.from_dlpack`.\n", + "\n", + "**Runs on:** any CUDA GPU." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. The kernel -- `@cute.kernel`\n", + "\n", + "A `@cute.kernel` function is the code **one GPU thread** runs; the GPU runs thousands of copies of it\n", + "in parallel. A thread locates its own slot from three built-ins -- `cute.arch.thread_idx()` (lane\n", + "within the block), `cute.arch.block_idx()` (which block), and `cute.arch.block_dim()` (block size) --\n", + "each a 3-tuple `(x, y, z)`; we use `x`. The global element index is `block * block_size + thread`.\n", + "\n", + "Kernel arguments are **`cutlass.Array`**: a typed handle to GPU memory you index like a Python\n", + "sequence. `cute.printf` prints from the device -- it takes a format string with `{}` holes followed\n", + "by the values. Here only thread 0 greets us, since thousands of threads printing would flood the\n", + "console." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def hello_kernel(a: cutlass.Array, b: cutlass.Array, c: cutlass.Array, N: cutlass.Int32):\n", + " tx, _, _ = cute.arch.thread_idx() # lane within the block\n", + " bx, _, _ = cute.arch.block_idx() # which block\n", + " bdx, _, _ = cute.arch.block_dim() # threads per block\n", + " i = bx * bdx + tx # this thread's global element\n", + "\n", + " # One thread says hello (every thread printing would flood the console).\n", + " if i == 0:\n", + " cute.printf(\"hello from the GPU -- block {}, thread {}\\n\", bx, tx)\n", + "\n", + " # Each thread adds one element. Guard the tail: the grid rounds up to whole\n", + " # blocks, so the last block can have threads whose index runs past the array.\n", + " if i < N:\n", + " c[i] = a[i] + b[i]" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 2. The launcher -- `@cute.jit`\n", + "\n", + "The `@cute.jit` host function picks the launch shape and starts the kernel with\n", + "`.launch(grid=, block=)`: `grid` is the number of blocks, `block` the threads per block. We use one\n", + "thread per element and round the grid up so every element is covered.\n", + "\n", + "To call it, wrap each PyTorch CUDA tensor with `cute.runtime.from_dlpack` -- a zero-copy view the DSL\n", + "accepts wherever a `cutlass.Array` is expected -- then just call the host function." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def hello(a: cutlass.Array, b: cutlass.Array, c: cutlass.Array, N: cutlass.Int32):\n", + " block = (256, 1, 1)\n", + " grid = ((N + 255) // 256, 1, 1) # one thread per element, rounded up to whole blocks\n", + " hello_kernel(a, b, c, N).launch(grid=grid, block=block)\n", + "\n", + "\n", + "N = 1024\n", + "a = torch.randn(N, dtype=torch.float32, device=\"cuda\")\n", + "b = torch.randn(N, dtype=torch.float32, device=\"cuda\")\n", + "c = torch.zeros(N, dtype=torch.float32, device=\"cuda\")\n", + "\n", + "# from_dlpack wraps each CUDA tensor as a cutlass.Array with no copy.\n", + "hello(cute.runtime.from_dlpack(a), cute.runtime.from_dlpack(b), cute.runtime.from_dlpack(c), N)\n", + "\n", + "torch.testing.assert_close(c.cpu(), (a + b).cpu(), atol=1e-5, rtol=1e-5)\n", + "print(\"PASS\")\n", + "\n", + "# Expected output (the GPU greeting prints once, from thread 0):\n", + "# hello from the GPU -- block 0, thread 0\n", + "# PASS" + ] + }, + { + "cell_type": "markdown", + "id": "7", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "- Change `N` and watch the tail guard `if i < N` earn its keep when it is not a multiple of 256.\n", + "- Make **every** thread print its index -- then see why the greeting is guarded to thread 0.\n", + "- Swap the `+` for a `*`, or add a third input.\n", + "\n", + "Next, `02_control_flow` covers loops and compile-time branching; over in `2_primitives`,\n", + "`01_array_concepts` digs into `cutlass.Array` itself -- the memory spaces, multi-dimensional\n", + "indexing, and vectorized slices." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/1_dsl_features/02_control_flow.ipynb b/examples/python/CuTeDSL/notebooks/1_dsl_features/02_control_flow.ipynb new file mode 100644 index 0000000000..1475e12ddf --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/1_dsl_features/02_control_flow.ipynb @@ -0,0 +1,422 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "8a35910d", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Control flow: meta loops vs staged loops\n", + "\n", + "A cutlass kernel is **traced** once on the host (the Python body runs to build IR) and then\n", + "**executed** many times on the GPU. Every loop you write lands on one side of that line:\n", + "\n", + "- A **meta loop** (`cutlass.range_constexpr`) is a *metaprogram* — it runs in the **Python\n", + " interpreter, on the host, while tracing**. The body is walked `N` times right there, emitting\n", + " `N` copies of the IR (fully unrolled). Because it is literally Python, the body can do\n", + " *anything Python can* — compute with `int`s, branch, even `raise`.\n", + "- A **staged loop** (`cutlass.range`) traces its body **once** into a single GPU `scf.for`\n", + " that the **hardware** iterates at run time. The loop index is a *staged* value, not a Python\n", + " `int` — there is no concrete number to hold yet.\n", + "\n", + "The cleanest way to *see* the split is a trace-time `print`: it fires once per host walk of the\n", + "body, so it counts iterations for you. This chapter needs no data or arrays — just loops and\n", + "prints." + ] + }, + { + "cell_type": "markdown", + "id": "3db9116b", + "metadata": {}, + "source": [ + "**You'll learn:** the **trace-time vs run-time** split that governs control flow; the meta loop\n", + "**`cutlass.range_constexpr(n)`** (body runs `n` times *in Python*, fully unrolled, bound must be\n", + "a Python `int`) versus the staged loop **`cutlass.range(n)`** (body traces *once* into a GPU\n", + "`scf.for`, bound may be a dynamic `Int32`); why a meta loop can run **arbitrary Python** (we make\n", + "one `raise` mid-trace); the partial-unroll option **`cutlass.range(n, unroll=K)`**; the\n", + "compile-time branch **`if cutlass.const_expr(cond):`** versus a plain runtime `if`; and how to\n", + "read all of it straight off `print()` (host, trace-time) vs `cute.printf()` (device, run-time).\n", + "\n", + "**Runs on:** any CUDA GPU." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "619ebcc9", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "\n", + "cutlass.cuda.initialize_cuda_context() # once per process, before launching\n", + "print(\"imports OK\")" + ] + }, + { + "cell_type": "markdown", + "id": "ab33336a", + "metadata": {}, + "source": [ + "## 1. Two kinds of loop\n", + "\n", + "The same job — walk the indices `0 .. N-1` — written two ways. The only difference is the loop\n", + "driver, and that single word decides **where the iteration happens**.\n", + "\n", + "| loop | what it is | where it iterates | bound | body trace count | emits |\n", + "|---|---|---|---|---|---|\n", + "| `cutlass.range_constexpr(N)` | **metaprogram** (Python) | **host, at trace time** | must be a Python `int` | **`N` times** | `N` straight-line copies (fully unrolled) |\n", + "| `cutlass.range(N)` | **staged** (GPU) | **device, at run time** | Python `int` *or* dynamic `Int32` | **once** | one `scf.for` the hardware loops |\n", + "\n", + "We prove the split with a trace-time `print` in each body. `print` is ordinary Python, so it\n", + "fires exactly as often as the host walks the body: **`N` lines** for the meta loop, **one** for\n", + "the staged loop. In the meta loop the index `i` is a real Python `int` (so `i` accumulates a\n", + "host-side total that gets *baked* into the kernel as a constant); in the staged loop `i` is a\n", + "staged `Int32` the GPU supplies each trip, so the work happens on the device." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "91aacb8e", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def loop_unrolled(N: cutlass.Constexpr):\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " if tx == 0:\n", + " total = 0 # a plain Python int -- accumulated at TRACE time\n", + " # N is a Constexpr (Python int), so this body runs N times right here, in Python.\n", + " # The print fires once per iteration; `total` is computed on the host and baked in.\n", + " for i in cutlass.range_constexpr(N):\n", + " print(f\"[trace] range_constexpr body, i = {i} (a Python int)\")\n", + " total += i\n", + " # `total` is now a constant -- cute.printf just emits that baked number.\n", + " cute.printf(\"[device] unrolled: sum of 0..{} baked at trace = {}\", N - 1, total)\n", + "\n", + "\n", + "@cute.jit\n", + "def loop_unrolled_host(N: cutlass.Constexpr):\n", + " loop_unrolled(N).launch(grid=(1, 1, 1), block=(1, 1, 1))" + ] + }, + { + "cell_type": "markdown", + "id": "a6dc1347", + "metadata": {}, + "source": [ + "The staged version differs in two words: `N` is a dynamic `cutlass.Int32` (the bound is decided\n", + "at *launch*, not when the kernel is written) and the loop is `cutlass.range`. Now the body\n", + "traces **once**, the accumulator is threaded through the GPU loop for you, and the sum is\n", + "computed on the **device**." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "de757977", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def loop_staged(N: cutlass.Int32):\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " if tx == 0:\n", + " total = cutlass.Int32(0)\n", + " # N is a dynamic Int32, so this is a real GPU loop: the body traces ONCE and the\n", + " # print fires ONCE. `i` is a staged Int32 -- the GPU supplies its value each trip.\n", + " for i in cutlass.range(N):\n", + " print(f\"[trace] range body, i is a staged {type(i).__name__} (traced once)\")\n", + " total = total + i\n", + " cute.printf(\"[device] staged: sum of 0..{} computed on the GPU = {}\", N - 1, total)\n", + "\n", + "\n", + "@cute.jit\n", + "def loop_staged_host(N: cutlass.Int32):\n", + " loop_staged(N).launch(grid=(1, 1, 1), block=(1, 1, 1))" + ] + }, + { + "cell_type": "markdown", + "id": "013132a9", + "metadata": {}, + "source": [ + "## 2. Run both and watch the trace\n", + "\n", + "No arrays, no data setup — just stage each kernel and read the prints. The **trace** prints fire\n", + "while each kernel is *compiled* (host, Python): `N` lines from the meta loop, **one** from the\n", + "staged loop. The **device** prints fire after the launch, once we sync; both report the same sum." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "89d3aa20", + "metadata": {}, + "outputs": [], + "source": [ + "N = 4\n", + "\n", + "print(\"--- staging loop_unrolled (range_constexpr) ---\")\n", + "loop_unrolled_host(N) # N trace lines fire HERE, during compile\n", + "\n", + "print(\"\\n--- staging loop_staged (range) ---\")\n", + "loop_staged_host(N) # exactly ONE trace line fires here\n", + "\n", + "cutlass.cuda.stream_sync(cutlass.cuda.default_stream()) # flush the device prints\n", + "\n", + "# Expected output:\n", + "# --- staging loop_unrolled (range_constexpr) ---\n", + "# [trace] range_constexpr body, i = 0 (a Python int)\n", + "# [trace] range_constexpr body, i = 1 (a Python int)\n", + "# [trace] range_constexpr body, i = 2 (a Python int)\n", + "# [trace] range_constexpr body, i = 3 (a Python int)\n", + "#\n", + "# --- staging loop_staged (range) ---\n", + "# [trace] range body, i is a staged Int32 (traced once) <- exactly ONE line\n", + "#\n", + "# [device] unrolled: sum of 0..3 baked at trace = 6\n", + "# [device] staged: sum of 0..3 computed on the GPU = 6" + ] + }, + { + "cell_type": "markdown", + "id": "0ae8376c", + "metadata": {}, + "source": [ + "## 3. The meta loop is just Python — so it can `raise`\n", + "\n", + "Because `range_constexpr` runs in the Python interpreter at trace time, the body can do anything\n", + "Python can: compute, branch on Python values, call helpers — even **raise an exception**. A\n", + "raise here aborts *compilation*; it never reaches the GPU. That's the whole point of a\n", + "metaprogram — you can validate and shape the kernel with ordinary Python before a single device\n", + "instruction exists. (Try this in a `range` body instead and there is no Python `i` to test — the\n", + "check would have to become real GPU code.)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "f03d6f03", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def checked_unroll(N: cutlass.Constexpr):\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " \n", + " for i in cutlass.range_constexpr(N):\n", + " # This runs in Python, WHILE TRACING. Ordinary Python control flow works --\n", + " # including raising, which stops compilation before any IR is emitted for i>=3.\n", + " if cutlass.const_expr(i == 3):\n", + " raise ValueError(\n", + " f\"meta-loop guard tripped at i={i}: this is a Python exception raised \"\n", + " \"during TRACING, before any GPU instruction is emitted\"\n", + " )\n", + " cute.printf(\"[device] checked i={}\", i)\n", + "\n", + "\n", + "@cute.jit\n", + "def checked_unroll_host(N: cutlass.Constexpr):\n", + " checked_unroll(N).launch(grid=(1, 1, 1), block=(1, 1, 1))" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6984a726", + "metadata": {}, + "outputs": [], + "source": [ + "print(\"N=3: the meta loop never reaches i==3, so it traces + runs cleanly:\")\n", + "checked_unroll_host(3)\n", + "cutlass.cuda.stream_sync(cutlass.cuda.default_stream())\n", + "\n", + "print(\"\\nN=8: the meta loop hits i==3 and raises -- DURING tracing, on the host:\")\n", + "try:\n", + " checked_unroll_host(8) # raises while the Python body is being walked\n", + "except ValueError as e:\n", + " print(f\" caught at trace time: {e}\")\n", + "\n", + "# Expected output:\n", + "# N=3: ... traces + runs cleanly:\n", + "# [device] checked i=0\n", + "# [device] checked i=1\n", + "# [device] checked i=2\n", + "#\n", + "# N=8: ... raises -- DURING tracing, on the host:\n", + "# caught at trace time: meta-loop guard tripped at i=3: this is a Python exception ..." + ] + }, + { + "cell_type": "markdown", + "id": "2fc5a83e", + "metadata": {}, + "source": [ + "## 4. Partial unroll: `range(N, unroll=K)`\n", + "\n", + "A staged loop can still be unrolled — just not all the way. `cutlass.range(N, unroll=K)` keeps a\n", + "real GPU loop (the bound stays dynamic, the body still traces **once**) but asks the compiler to\n", + "emit `K` copies of the body per trip, so the loop runs `N/K` times with `K`-wide bodies. It is\n", + "the middle ground: fewer loop-control instructions and more instruction-level parallelism than\n", + "the default, without the code-size blow-up of a full unroll. The result is identical — only the\n", + "generated loop shape changes. (`unroll=0`/`1` is the plain staged loop; `unroll_full=True` would\n", + "flatten a dynamic loop completely.)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "d45b4382", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def loop_unroll4(N: cutlass.Int32):\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " if tx == 0:\n", + " total = cutlass.Int32(0)\n", + " # Still a GPU loop (N is dynamic) -- the body traces once; the COMPILER emits 4\n", + " # bodies per trip. So the trace print still fires exactly once.\n", + " for i in cutlass.range(N, unroll=4):\n", + " print(f\"[trace] unroll=4 body, traced once (i is staged {type(i).__name__})\")\n", + " total = total + i\n", + " cute.printf(\"[device] unroll=4: sum of 0..{} = {}\", N - 1, total)\n", + "\n", + "\n", + "@cute.jit\n", + "def loop_unroll4_host(N: cutlass.Int32):\n", + " loop_unroll4(N).launch(grid=(1, 1, 1), block=(1, 1, 1))\n", + "\n", + "\n", + "loop_unroll4_host(N)\n", + "cutlass.cuda.stream_sync(cutlass.cuda.default_stream())\n", + "\n", + "# Expected output:\n", + "# [trace] unroll=4 body, traced once (i is staged Int32) <- still ONE trace line\n", + "# [device] unroll=4: sum of 0..3 = 6" + ] + }, + { + "cell_type": "markdown", + "id": "51779a4c", + "metadata": {}, + "source": [ + "## 5. Two kinds of branch\n", + "\n", + "Branches split the same way as loops:\n", + "\n", + "- `if cutlass.const_expr(cond):` — `cond` is a Python value fixed at trace time, so the host\n", + " picks **one** side and compiles only that; the other is never traced. A trace-time `print` in\n", + " each side proves it — only the taken side fires.\n", + "- a plain `if :` — the condition is a per-thread device value, so **both** sides compile\n", + " and the GPU chooses per thread (a real `scf.if`). You still write a plain Python `if`; the DSL\n", + " stages it because the condition is dynamic.\n", + "\n", + "No arrays here either — we branch on a `Constexpr` flag and on the per-thread `tidx`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "eaecdfa0", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def branch_demo(USE_DOUBLE: cutlass.Constexpr):\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " # COMPILE-TIME branch: USE_DOUBLE is a Python bool, so exactly ONE side is traced and\n", + " # emitted -- the other side's print never fires. `factor` is baked from the taken side.\n", + " if cutlass.const_expr(USE_DOUBLE):\n", + " print(\"[trace] const_expr branch: doubling side compiled in\")\n", + " factor = 2\n", + " else:\n", + " print(\"[trace] const_expr branch: identity side compiled in\")\n", + " factor = 1\n", + " # RUNTIME branch: tx is a staged value, so BOTH sides compile and the GPU picks per\n", + " # thread. Plain Python if -- the DSL stages it into a real scf.if.\n", + " if tx == 0:\n", + " cute.printf(\"[device] thread {} -> first lane (baked factor = {})\", tx, factor)\n", + " else:\n", + " cute.printf(\"[device] thread {} -> other lane (baked factor = {})\", tx, factor)\n", + "\n", + "\n", + "@cute.jit\n", + "def branch_demo_host(USE_DOUBLE: cutlass.Constexpr):\n", + " branch_demo(USE_DOUBLE).launch(grid=(1, 1, 1), block=(2, 1, 1))" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "247cb6c9", + "metadata": {}, + "outputs": [], + "source": [ + "for use_double in (True, False):\n", + " print(f\"--- staging branch_demo(USE_DOUBLE={use_double}) ---\")\n", + " branch_demo_host(use_double) # only the TAKEN const_expr side prints at trace\n", + " cutlass.cuda.stream_sync(cutlass.cuda.default_stream())\n", + " print()\n", + "\n", + "# Expected output (only ONE const_expr print per staging -- the taken side):\n", + "# --- staging branch_demo(USE_DOUBLE=True) ---\n", + "# [trace] const_expr branch: doubling side compiled in\n", + "# [device] thread 0 -> first lane (baked factor = 2)\n", + "# [device] thread 1 -> other lane (baked factor = 2)\n", + "#\n", + "# --- staging branch_demo(USE_DOUBLE=False) ---\n", + "# [trace] const_expr branch: identity side compiled in\n", + "# [device] thread 0 -> first lane (baked factor = 1)\n", + "# [device] thread 1 -> other lane (baked factor = 1)" + ] + }, + { + "cell_type": "markdown", + "id": "53d856c4", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. **Count the trace prints.** Bump `N` to 8 and re-run section 2: the meta loop now prints 8\n", + " trace lines, the staged loop still prints exactly one.\n", + "2. **Break the constexpr bound.** Change `loop_unrolled`'s `N` from `cutlass.Constexpr` to\n", + " `cutlass.Int32` and pass a runtime value — `range_constexpr` raises, because its bound must\n", + " be a Python `int` known at trace time. Switching that loop to `cutlass.range(N)` fixes it.\n", + "3. **Move the guard into a staged loop.** Put the `if i == 3: raise ...` from section 3 inside a\n", + " `cutlass.range(N)` loop. It does *not* raise per-trip — `i` is a staged value, so `i == 3` is\n", + " a device comparison, not a Python one; the meta-only trick is gone.\n", + "4. **Make the const_expr branch runtime.** In section 5, drop `cutlass.const_expr(...)` and pass\n", + " `USE_DOUBLE` as a `cutlass.Int32` flag compared `> 0` — now *both* sides compile (both trace\n", + " prints fire) and the GPU chooses at run time, just like the `tx == 0` guard already does." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/1_dsl_features/03_diagnostics.ipynb b/examples/python/CuTeDSL/notebooks/1_dsl_features/03_diagnostics.ipynb new file mode 100644 index 0000000000..cb1680389f --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/1_dsl_features/03_diagnostics.ipynb @@ -0,0 +1,312 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Compiler diagnostics: `warnings{...}` and `remarks{...}`\n", + "\n", + "The compiler does more than translate your kernel to PTX -- it can also *inspect* it and tell you\n", + "when something looks wrong. Two opt-in options turn those inspections on:\n", + "\n", + "- **`warnings{...}`** surfaces **legal-but-questionable** patterns: code that compiles and may run, but\n", + " has a real chance of hanging, faulting, or behaving differently than you intended.\n", + "- **`remarks{...}`** surfaces **perf-only** findings: the program is functionally correct, just slower\n", + " than it could be (for example, register spills into local memory).\n", + "\n", + "Both are **non-fatal** -- compilation still succeeds and you still get a kernel. A third severity,\n", + "the **error**, is *always on* and *always fatal*: it fires on a proven defect, with or without any\n", + "option. This chapter shows one of each, end to end, and exactly how to switch the inspections on." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** how to enable the compiler's diagnostic inspections with\n", + "`cute.compile(..., options=\"warnings{nvvm}\")` / `\"remarks{ptx}\"` (or the matching `CUTE_DSL_COMPILER_OPT`\n", + "environment variable); the three diagnostic **severities** -- *remark* (perf), *warning*\n", + "(questionable), *error* (defect) -- and which are fatal; and how to catch a fatal diagnostic in Python with `CompilerDiagnosticError`.\n", + "\n", + "**Runs on:** any CUDA GPU (the synchronization examples target sm_90+; the ptxas remark works on any\n", + "target). **Prereq:** Chapter on `cutlass.Array` and the memory spaces." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "from cutlass.experimental import primitives as prims \n", + "from cutlass.base_dsl.compiler import CompilerDiagnosticError" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## Levels and categories\n", + "\n", + "You turn diagnostics on by passing an **options string** to `cute.compile`.\n", + "The string has two axes: the severity level to show and the diagnostic category to collect.\n", + "Warnings and remarks are opt-in and non-fatal. Errors reported by an enabled category are always\n", + "shown and fail compilation; there is no separate `errors{...}` option.\n", + "\n", + "| Level | Enable with | Useful for | Fatal? |\n", + "|---|---|---|---|\n", + "| Info (remark) | `\"remarks\"` or `\"remarks{}\"` | Performance-only findings, including synchronization opportunities and ptxas resource reports. | No |\n", + "| Warning | `\"warnings\"` or `\"warnings{}\"` | Legal but questionable patterns that can hang, fault, or behave differently than intended. | No |\n", + "| Error | Enable the relevant category with `\"warnings{}\"` or `\"remarks{}\"`. | Proven defects reported by an enabled diagnostic category. | Yes |\n", + "\n", + "| Category | Enable with | Source | Useful levels |\n", + "|---|---|---|---|\n", + "| `nvvm` | `\"warnings{nvvm}\"`, `\"remarks{nvvm}\"` | NVVM-level primitive protocol diagnostics for operations such as `mbarrier`, bulk copy, TMA multicast, and `tcgen05`. | Error, warning, info (remark) |\n", + "| `ptxas` (selector: `ptx`) | `\"remarks{ptx}\"` | ptxas resource diagnostics surfaced through the remark stream, including register spills and local-memory usage. | Info (remark) |\n", + "\n", + "The exact same strings work as an environment variable, handy for turning diagnostics on without\n", + "editing code:\n", + "\n", + "```bash\n", + "CUTE_DSL_COMPILER_OPT=\"warnings{nvvm}\" python my_kernel.py\n", + "```\n", + "\n", + "That is the whole mechanism. This chapter is about *how to switch the inspections on* -- not about\n", + "the specific findings -- so the three examples below show one of each severity end to end." + ] + }, + { + "cell_type": "markdown", + "id": "4", + "metadata": {}, + "source": [ + "## 1. A warning -- legal but questionable\n", + "\n", + "A *warning* fires on code that the compiler will happily build, but that carries a real risk at run\n", + "time. Here the kernel guards a uniform bulk copy with a full-mask `elect.sync` -- a correct\n", + "single-issuer idiom -- but the launch uses a **partial warp** of 4 threads. A full-mask elect in a\n", + "trailing partial warp can hang or fault, so the compiler flags it. It is *questionable*, not\n", + "*proven wrong*, so it is a **warning**: compilation still succeeds." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def partial_warp_bulk_copy_kernel(gmem: cutlass.Array):\n", + " # A single-issuer bulk copy: one elected lane drives the cp.async.bulk for the whole warp.\n", + " mbar = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8)\n", + " smem_dst = cutlass.Array(cutlass.Int32, 4, space=cutlass.AddressSpace.smem, alignment=16)\n", + "\n", + " if prims.elect_sync(): # full-mask elect: one lane of the (assumed full) warp\n", + " prims.cp_async_bulk_shared_cluster_global(smem_dst, gmem, mbar, 16)\n", + "\n", + "\n", + "@cute.jit\n", + "def warning_host(gmem: cutlass.Array):\n", + " # block=(4,1,1) is only a PARTIAL warp -> the elect.sync above becomes questionable.\n", + " partial_warp_bulk_copy_kernel(gmem).launch(grid=(1, 1, 1), block=(4, 1, 1))\n", + "\n", + "\n", + "# A fake (compile-only) array + stream let us compile without launching on a real GPU.\n", + "gmem = cutlass.runtime.make_fake_compact_array(cutlass.Int32, (4,), assumed_align=16)\n", + "\n", + "print(\">>> compiling WITHOUT warnings{nvvm} (no diagnostic):\")\n", + "cute.compile(warning_host, gmem)\n", + "print(\" compiled, silent.\\n\")\n", + "\n", + "print(\">>> compiling WITH options='warnings{nvvm}':\")\n", + "cute.compile(warning_host, gmem, options=\"warnings{nvvm}\")\n", + "print(\" compiled anyway -- the warning is non-fatal.\")\n", + "\n", + "# Expected output: the first compile is silent; the second prints a source-located\n", + "# warning before succeeding (compilation does NOT fail):\n", + "#\n", + "# warning[nvvm-diag:C13]: full-mask elect.sync guards ... in a partial-warp launch block\n", + "# --> .../03_diagnostics...:NN:7\n", + "# ...\n", + "# suggestion: launch with a warp-multiple block size ..." + ] + }, + { + "cell_type": "markdown", + "id": "6", + "metadata": {}, + "source": [ + "## 2. An error -- a proven defect, always fatal\n", + "\n", + "An *error* fires on a defect the compiler can prove, and it **always** fails compilation. There is no\n", + "\"errors\" flag: the compiler's checkers run whenever you ask for `warnings{nvvm}` (or `remarks{nvvm}`), and\n", + "once they run, any error they find is fatal regardless of which severity you asked to *see*.\n", + "\n", + "This kernel initializes an `mbarrier` to expect **exactly one** arrival per phase, then lets *every*\n", + "thread arrive on it -- unguarded. That over-arrives the barrier and flips its phase unexpectedly: a\n", + "real bug, not a style nit. Because it is fatal, we wrap the compile in `try/except` and catch\n", + "`CompilerDiagnosticError`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "7", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def unguarded_arrive_kernel():\n", + " mbar = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem, alignment=8)\n", + "\n", + " if prims.elect_sync():\n", + " prims.mbarrier_init(mbar, 1) # expect exactly ONE arrival per phase\n", + " prims.fence_mbarrier_init()\n", + " cute.arch.barrier()\n", + "\n", + " # BUG: no single-issuer guard, so all 32 threads arrive on a count=1 barrier.\n", + " prims.mbarrier_arrive(mbar, count=1)\n", + " prims.mbarrier_try_wait_parity(mbar, 0, time_limit=10_000_000)\n", + "\n", + "\n", + "@cute.jit\n", + "def error_host():\n", + " unguarded_arrive_kernel().launch(grid=(1, 1, 1), block=(32, 1, 1))\n", + "\n", + "\n", + "print(\">>> compiling WITH options='warnings{nvvm}' (runs the checkers):\")\n", + "try:\n", + " # warnings{nvvm} runs the checkers; any error they find is fatal no matter which severity you ask for.\n", + " cute.compile(error_host, options=\"warnings{nvvm}\")\n", + " print(\" unexpected: compile succeeded\")\n", + "except CompilerDiagnosticError as exc:\n", + " print(\" compilation FAILED, as it should -- here is the diagnostic:\\n\")\n", + " print(exc)" + ] + }, + { + "cell_type": "markdown", + "id": "8", + "metadata": {}, + "source": [ + "## 3. A remark -- correct, just slower\n", + "\n", + "A *remark* is perf-only: the program is functionally correct, but the compiler noticed something that\n", + "costs performance. The classic example is a **register spill** -- when a kernel keeps too many values\n", + "live at once, `ptxas` runs out of registers and parks the overflow in local memory, adding traffic\n", + "and latency.\n", + "\n", + "This kernel loads a 64-element window into registers and writes it back **reversed**, which keeps all\n", + "64 values live simultaneously. We compile it with a deliberately tight register budget\n", + "(`--maxrregcount=24`) to force the spill, then ask for the `ptx` remark category. Leave\n", + "`--remark-output` unset so the Python diagnostic renderer prints source frames. Because `ptxas`\n", + "reports spill totals at kernel granularity, the source frame marks the reported kernel / likely\n", + "pressure region rather than an exact spill instruction." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9", + "metadata": {}, + "outputs": [], + "source": [ + "_N = 64 # values per thread -- a full reversed window kept live forces register pressure\n", + "\n", + "\n", + "@cute.kernel\n", + "def reverse_store_kernel(src: cutlass.Array, dst: cutlass.Array):\n", + " tid, _, _ = cute.arch.thread_idx()\n", + " base = tid * cutlass.Int32(_N)\n", + " window = src[base:_N] # load the whole V-wide window into registers\n", + " # Storing it reversed keeps every element of `window` live at once -> high register pressure.\n", + " for i in cutlass.range_constexpr(_N):\n", + " dst[base + cutlass.Int32(i)] = window[_N - 1 - i]\n", + "\n", + "\n", + "@cute.jit\n", + "def remark_host(src: cutlass.Array, dst: cutlass.Array):\n", + " reverse_store_kernel(src, dst).launch(grid=(1, 1, 1), block=(32, 1, 1))\n", + "\n", + "\n", + "n = _N * 32\n", + "src = cutlass.runtime.make_fake_compact_array(cutlass.Int32, (n,), assumed_align=16)\n", + "dst = cutlass.runtime.make_fake_compact_array(cutlass.Int32, (n,), assumed_align=16)\n", + "\n", + "print(\">>> compiling WITH the ptx remark category + a tight register budget:\")\n", + "# remarks{ptx} asks ptxas for perf remarks; --ptxas-options squeezes the register budget\n", + "# so the spill actually happens.\n", + "cute.compile(\n", + " remark_host,\n", + " src,\n", + " dst,\n", + " options=(\n", + " \"remarks{ptx} \"\n", + " \"--ptxas-options '--maxrregcount=24 --override-directive-values'\"\n", + " ),\n", + ")\n", + "print(\" compiled (the spill is perf-only, never fatal). See remark[ptxas] above.\")\n", + "\n", + "# Expected output: compilation succeeds and prints a ptxas spill remark with a source frame." + ] + }, + { + "cell_type": "markdown", + "id": "10", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. **Silence vs. surface.** Re-run section 1 with `options=\"remarks{nvvm}\"` instead of `\"warnings{nvvm}\"`:\n", + " the partial-warp *warning* no longer prints (it is gated to `warnings{nvvm}`), yet the compile still\n", + " succeeds. Severity options control what you *see*, not what the compiler *checks*.\n", + "2. **The error is unconditional.** In section 2, drop the `options=` argument entirely. The\n", + " over-arrive is a *defect*, so... it depends: with no diagnostic option the sync checker never runs, so the\n", + " compile *succeeds silently*. Add `options=\"remarks{nvvm}\"` and the same defect now fails the\n", + " compile -- proof that asking for diagnostics (either severity) is what runs the checkers, after\n", + " which the error is always fatal.\n", + "3. **Fix the bug.** Guard the arrive in section 2 with `if prims.elect_sync():` (or `if tid == 0:`,\n", + " reading `tid` from `cute.arch.thread_idx()`) so a single thread arrives. Re-compile with\n", + " `options=\"warnings{nvvm}\"` -- the error is gone.\n", + "4. **Make the spill worse (or vanish).** In section 3, raise `_N` to 96 and watch the spill counts\n", + " in the report climb; or drop the `--ptxas-options ...` part so ptxas keeps its full register budget,\n", + " and the spill remark disappears -- there was nothing to spill.\n", + "5. **Use the environment variable.** Instead of `options=\"...\"`, set\n", + " `CUTE_DSL_COMPILER_OPT=\"warnings{nvvm}\"` in your shell and run a plain `cute.compile(...)` with no\n", + " options string -- same diagnostics, no code change." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/1_dsl_features/04_zero_cost_abstraction.ipynb b/examples/python/CuTeDSL/notebooks/1_dsl_features/04_zero_cost_abstraction.ipynb new file mode 100644 index 0000000000..72fdac40a2 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/1_dsl_features/04_zero_cost_abstraction.ipynb @@ -0,0 +1,254 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Zero-cost abstraction: classes and polymorphism compile away\n", + "\n", + "A `@cute.kernel` body is **Python that runs at trace time** (the metaprogram idea from\n", + "*02_control_flow*). So the abstractions you reach for — classes, objects, methods, even\n", + "*polymorphism* — are resolved by the Python interpreter while the kernel is built, and compile\n", + "down to **nothing extra**. The device program is exactly the arithmetic and memory ops you emit.\n", + "We show two examples and read the generated PTX to prove the abstractions are free." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "**You'll learn:** that ordinary Python **classes / objects / methods** used inside a kernel are\n", + "**zero-cost** (resolved at trace time, never present on the device); that a **polymorphic pipeline**\n", + "— a list of objects sharing a (virtual) method — *unrolls and inlines* into straight-line fused\n", + "arithmetic; and how to read the generated **PTX** (`cute.compile[KeepPTX(True)](...).__ptx__`) to\n", + "confirm there is no object, no vtable, and no dispatch.\n", + "\n", + "**Runs on:** any CUDA GPU." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "from cutlass.cute import KeepPTX # lets us read the generated PTX back via __ptx__\n", + "from typing import Any\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 1. A class is free\n", + "\n", + "`Printer` is an ordinary Python class. Inside the kernel we build one and call its `.print(val)`\n", + "method **twice** — once naming the class directly, once through a class handed in as a `Constexpr`\n", + "argument. Both are pure Python, run while the host traces the body; the only thing *emitted* is the\n", + "`cute.printf(val)` inside `.print`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "class Printer:\n", + " def __init__(self, printer_method):\n", + " self.printer_method = printer_method\n", + "\n", + " def print(self, v):\n", + " self.printer_method(v)\n", + "\n", + "\n", + "@cute.kernel\n", + "def printer_kernel(class_type: cutlass.Constexpr[Any], val: cutlass.Int32):\n", + " pm = cute.printf\n", + " Printer(pm).print(val) # construct a class, call its method\n", + " class_type(pm).print(val) # same, via a class passed in as Constexpr\n", + "\n", + "\n", + "@cute.jit\n", + "def printer(class_type: cutlass.Constexpr[Any], val: cutlass.Int32):\n", + " printer_kernel(class_type, val).launch(grid=[1, 1, 1], block=[1, 1, 1])" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "Run it, then read the PTX. The whole device program is **two `vprintf` calls** (one per `.print`).\n", + "`Printer` survives only inside the *mangled kernel name* — it was a compile-time `Constexpr` baked\n", + "into the kernel's identity, never an object on the device." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "cutlass.cuda.initialize_cuda_context()\n", + "printer(Printer, 1)\n", + "cutlass.cuda.stream_sync(cutlass.cuda.default_stream()) # -> prints 1 twice\n", + "\n", + "compiled = cute.compile[KeepPTX(True)](printer, Printer, 1)\n", + "ptx = compiled.__ptx__\n", + "print(ptx)\n", + "print(\"\\n>>> vprintf call sites:\", sum(1 for l in ptx.splitlines() if 'call' in l and 'vprintf' in l))\n", + "\n", + "# Expected: the device program is just\n", + "# .visible .entry kernel_cutlass_..._class___main__Printer__0( ... ) <- 'Printer' only in the baked name\n", + "# call.uni (retval0), vprintf, (...); <- .print(val) #1\n", + "# call.uni (retval0), vprintf, (...); <- .print(val) #2\n", + "# >>> vprintf call sites: 2\n", + "# 1\n", + "# 1" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 2. Polymorphism is free\n", + "\n", + "Now the real thing: a chain of elementwise ops, written as a small class hierarchy with a *virtual*\n", + "`apply()` — a base `Op` and `AddN` / `MulN` overriding it. The kernel walks a **build-time list** of\n", + "op objects and calls `op.apply(x)` on each. Both the list and the `for` loop are metaprogramming:\n", + "the loop unrolls at trace time and each virtual `apply()` resolves in Python and inlines — leaving\n", + "straight-line fused arithmetic, with no object, no vtable, no dispatch." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "class Op:\n", + " def apply(self, x):\n", + " raise NotImplementedError # 'virtual' -- overridden below\n", + "\n", + "\n", + "class AddN(Op):\n", + " def __init__(self, n):\n", + " self.n = n\n", + "\n", + " def apply(self, x):\n", + " return x + self.n\n", + "\n", + "\n", + "class MulN(Op):\n", + " def __init__(self, k):\n", + " self.k = k\n", + "\n", + " def apply(self, x):\n", + " return x * self.k\n", + "\n", + "\n", + "@cute.kernel\n", + "def fused_kernel(arr: cutlass.Array, ops: cutlass.Constexpr[Any]):\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " x = arr[tx]\n", + " for op in ops: # build-time list -> loop UNROLLS at trace time\n", + " x = op.apply(x) # virtual call resolves in Python -> op INLINES\n", + " arr[tx] = x # -> straight-line fused arithmetic in the IR\n", + "\n", + "\n", + "@cute.jit\n", + "def run(arr: cutlass.Array, ops: cutlass.Constexpr[Any]):\n", + " fused_kernel(arr, ops).launch(grid=[1, 1, 1], block=[256, 1, 1])" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "Build the pipeline with ordinary Python — a list of polymorphic ops — run it, and check the result.\n", + "Then read the fused kernel's PTX: the three `apply()` calls have become a couple of arithmetic\n", + "instructions with **zero `call` sites**. Swap or extend the op list and the kernel respecializes for\n", + "free." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "ops = [AddN(5.0), MulN(2.0), AddN(-1.0)] # (x + 5) * 2 - 1, built from objects\n", + "n = 256\n", + "data = torch.arange(n, dtype=torch.float32).cuda()\n", + "run(cute.runtime.from_dlpack(data), ops)\n", + "cutlass.cuda.stream_sync(cutlass.cuda.default_stream())\n", + "ref = (torch.arange(n, dtype=torch.float32).cuda() + 5.0) * 2.0 - 1.0\n", + "torch.testing.assert_close(data, ref, atol=1e-3, rtol=1e-3)\n", + "print(\"pipeline PASS\")\n", + "\n", + "# Read the fused kernel's PTX -- the op pipeline collapsed to straight-line arithmetic.\n", + "data2 = torch.arange(n, dtype=torch.float32).cuda()\n", + "ptx = cute.compile[KeepPTX(True)](run, cute.runtime.from_dlpack(data2), ops).__ptx__\n", + "for l in ptx.splitlines():\n", + " if any(k in l for k in ('ld.global', 'add.f32', 'mul.f32', 'fma.', 'st.global')):\n", + " print(\" \", l.strip())\n", + "print(\">>> call sites in the fused kernel:\", ptx.count('call'))\n", + "\n", + "# Expected output: three ops -> two arithmetic instructions, no calls:\n", + "# ld.global.b32 %r2, [%rd3];\n", + "# add.f32 %r3, %r2, 0f40A00000; <- + 5.0\n", + "# fma.rn.f32 %r4, %r3, 0f40000000, 0fBF800000; <- * 2.0 - 1.0 (fused)\n", + "# st.global.b32 [%rd3], %r4;\n", + "# >>> call sites in the fused kernel: 0\n", + "# pipeline PASS" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Takeaway\n", + "\n", + "Three polymorphic `apply()` calls compiled to **two arithmetic instructions and zero `call`s**. The\n", + "`Op` hierarchy, the virtual dispatch, the Python `for` over the op list — all of it ran in the Python\n", + "interpreter at trace time and left no trace on the device. That is *zero-cost abstraction*: you write\n", + "the kernel with the full Python object model, and pay for none of it at runtime. (Same mechanism as\n", + "the meta loop in *02_control_flow* — the kernel body is a metaprogram — applied to objects.)" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. **Respecialize for free.** Change `ops` to `[MulN(3.0), AddN(7.0)]` and re-run — no kernel edit,\n", + " just a different build-time list, and the PTX changes to match.\n", + "2. **Add an op.** Define `class Square(Op)` whose `apply` returns `x * x`, drop it into the list, and\n", + " re-read the PTX: one more `mul.f32`, still **zero `call`s**.\n", + "3. **Watch the loop unroll.** Put a trace-time `print(type(op).__name__)` in the `for` loop — it\n", + " fires once per op, in Python, while tracing (exactly like the meta loop in *02_control_flow*).\n", + "4. **Other artifacts.** `compiled.__sass__` shows the final SASS the same way `__ptx__` shows PTX." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/2_primitives/01_array_concepts.ipynb b/examples/python/CuTeDSL/notebooks/2_primitives/01_array_concepts.ipynb new file mode 100644 index 0000000000..2652d67af2 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/2_primitives/01_array_concepts.ipynb @@ -0,0 +1,511 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Array concepts: `cutlass.Array` and the GPU memory spaces\n", + "\n", + "`cutlass.Array` is the workhorse memory type of this course: a typed handle to GPU memory that\n", + "knows its **shape**, **strides**, **dtype**, and **memory space**. One handle serves scalar `a[i]`,\n", + "multi-dimensional `a[r, c]`, and offset `subview` access — no hand-written pointer arithmetic.\n", + "\n", + "The other half is **where** the memory lives. One `Array` type covers four GPU memory spaces, and\n", + "the space decides **who owns it, how long it lives, and which threads see it**. Two kernels carry\n", + "the chapter: one allocates arrays **inside** the kernel, the other uses arrays declared **outside**\n", + "it at module scope. Each prints the arrays it touches." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** the central data type of the DSL, **`cutlass.Array`**. Kernel arguments are\n", + "Arrays; the **slice = vectorized load/store** idiom (`a[idx:V]` moves `V` contiguous elements in\n", + "one wide transaction once you promise alignment); how to read an Array's printed layout (dtype,\n", + "space, shape, alignment); the **local / shared / global / constant** memory spaces and their\n", + "ownership, lifetime, and visibility; allocating **local** and **shared** scratch inside a kernel\n", + "(with the barrier that makes shared safe); declaring **global** and **constant** statics at module\n", + "scope; and `subview`. We build a complete vectorized vector-add along the way.\n", + "\n", + "**Runs on:** any CUDA GPU. (Registers as *values* are `cutlass.Vector`, in `02_vector_concepts`;\n", + "the `cute.Tensor` -> `cutlass.Array` bridge is `03_cute_interop`.)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## The four spaces\n", + "\n", + "`cutlass.Array(dtype, size, space=...)` picks its allocator from the **space**. Each maps to a\n", + "familiar CUDA C++ declaration:\n", + "\n", + "| `cutlass.Array(...)` | CUDA C++ analog | space | who sees it / lifetime |\n", + "|---|---|---|---|\n", + "| `cutlass.Array(f32, N)` | `float a[N];` | **local** (default) | one private copy per **thread** |\n", + "| `cutlass.Array(f32, N, space=…smem)` | `__shared__ float a[N];` | **shared** | one copy per **block** |\n", + "| `cutlass.Array(f32, N, space=…gmem, name=\"g\")` | `__device__ float g[N];` | **global** | one **static** buffer; all threads + launches share it |\n", + "| `cutlass.Array(f32, N, space=…cmem, name=\"c\", init=…)` | `__constant__ float c[N];` | **constant** | read-only static, broadcast to all threads |\n", + "\n", + "**local** and **shared** are scratch you allocate **inside** a kernel; **global** and **constant**\n", + "are **statics** you declare at module scope and use by name. `size` is always a compile-time\n", + "constant — these are static buffers, not a per-thread `malloc`. Global memory also arrives the\n", + "everyday way, as a **kernel argument** (a tensor via `from_dlpack`)." + ] + }, + { + "cell_type": "markdown", + "id": "4", + "metadata": {}, + "source": [ + "## 1. One thread, a vector of elements\n", + "\n", + "The kernel is three lines, so the only design decision is *how much work each thread does*. We give\n", + "every thread a window of `V` **contiguous** elements and let it process the whole window at once.\n", + "\n", + "The mechanism is the slice. `a[idx:V]` reads `V` elements starting at `idx`; assigning to `c[idx:V]`\n", + "writes them back. The second number is a **count**, not a Python stop index -- `a[idx:V]` means \"`V`\n", + "elements from `idx`\", not \"elements `idx` up to `V`\". So a thread owning the span `[idx, idx+V)`\n", + "loads both inputs, adds them register-to-register, and stores the result.\n", + "\n", + "Whether that slice becomes **one wide transaction** or `V` scalar ones comes down to **alignment**. A\n", + "128-bit load of four `f32`s is only legal from a 16-byte-aligned address, and the compiler will not\n", + "assume that on its own. We make the promise when we wrap the buffer below\n", + "(`from_dlpack(a, assumed_align=16)`), and it genuinely holds: every thread's window starts at\n", + "`idx = (block * block_size + thread) * V`, a multiple of `V = 4` elements, i.e. 16 bytes.\n", + "\n", + "| | one 128-bit transaction | `V` scalar transactions |\n", + "|---|---|---|\n", + "| when | `assumed_align=16` promised | no alignment promise |\n", + "| `a[idx:V]` lowers to | `ld.global.v4.b32` / `v2.b64` | four `ld.global.b32` |\n", + "| per slice | **1** instruction | **4** instructions |\n", + "\n", + "The exact wide opcode is GPU-specific -- `ld.global.v4.b32` (four 32-bit words) on Ampere/Hopper, or\n", + "`ld.global.v2.b64` (two 64-bit words) on Blackwell, which packs the four values so one `add.f32x2`\n", + "adds two at a time. Same 128 bits either way; section 4 prints whichever your GPU uses." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def vector_add_kernel(\n", + " a: cutlass.Array, b: cutlass.Array, c: cutlass.Array, N: cutlass.Int32, V: cutlass.Constexpr\n", + "):\n", + " # Step 1. Map this thread to the first element of its V-wide window. The global\n", + " # lane is (block * block_size + thread); scaling by V gives the window's start.\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " bx, _, _ = cute.arch.block_idx()\n", + " bdx, _, _ = cute.arch.block_dim()\n", + " idx = (bx * bdx + tx) * V\n", + "\n", + " # Step 2. Guard the tail, then add one V-wide window. The grid rounds up to whole\n", + " # blocks, so the last block has threads whose window starts past the array; skip them.\n", + " if idx < N:\n", + " c[idx:V] = a[idx:V] + b[idx:V]" + ] + }, + { + "cell_type": "markdown", + "id": "6", + "metadata": {}, + "source": [ + "## 2. Launch: one thread per window\n", + "\n", + "Each thread handles `V` elements, so we need `N / V` threads. We pick a 1-D block of 256 threads (a\n", + "conventional, occupancy-friendly size) and round the grid up so no element is left uncovered -- the\n", + "in-kernel bounds check absorbs the few extra threads. We take `N` to be a multiple of `V`, so every\n", + "active thread's window lands fully inside the array." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "7", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def vector_add(\n", + " a: cutlass.Array, b: cutlass.Array, c: cutlass.Array, N: cutlass.Int32, V: cutlass.Constexpr\n", + "):\n", + " block = (256, 1, 1)\n", + " threads = N // V # one thread per V-wide window\n", + " grid = ((threads + block[0] - 1) // block[0], 1, 1) # round up to whole blocks\n", + "\n", + " vector_add_kernel(a, b, c, N, V).launch(grid=grid, block=block)" + ] + }, + { + "cell_type": "markdown", + "id": "8", + "metadata": {}, + "source": [ + "## 3. Run it and check against PyTorch\n", + "\n", + "`cutlass.Array` parameters are host-entry types, so there is no manual buffer setup: hand the PyTorch\n", + "CUDA tensors straight to the kernel through `cute.runtime.from_dlpack`, which wraps the existing GPU\n", + "memory with zero copies. We then check the result against PyTorch's own `a + b`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9", + "metadata": {}, + "outputs": [], + "source": [ + "N, V = 1 << 20, 4 # 1,048,576 elements, 4 per thread\n", + "\n", + "# The grid is sized as N // V, so V must divide N exactly. Otherwise the trailing\n", + "# elements would have no thread assigned, and a thread's window could run past the end.\n", + "assert N % V == 0, \"N must be a multiple of the vector width V\"\n", + "\n", + "a = torch.randn(N, dtype=torch.float32, device=\"cuda\")\n", + "b = torch.randn(N, dtype=torch.float32, device=\"cuda\")\n", + "c = torch.zeros(N, dtype=torch.float32, device=\"cuda\")\n", + "\n", + "# from_dlpack wraps each CUDA tensor as a cutlass.Array with no copy. assumed_align=16\n", + "# promises 16-byte-aligned buffers (torch's allocator guarantees far more), which is what\n", + "# lets the V=4 slice load/store become single 128-bit transactions (see section 4).\n", + "a_ = cute.runtime.from_dlpack(a, assumed_align=16)\n", + "b_ = cute.runtime.from_dlpack(b, assumed_align=16)\n", + "c_ = cute.runtime.from_dlpack(c, assumed_align=16)\n", + "\n", + "vector_add(a_, b_, c_, N, V)\n", + "\n", + "torch.testing.assert_close(c.cpu(), (a + b).cpu(), atol=1e-5, rtol=1e-5)\n", + "print(\"PASS\")\n", + "\n", + "# Expected output:\n", + "# PASS" + ] + }, + { + "cell_type": "markdown", + "id": "10", + "metadata": {}, + "source": [ + "## 4. Prove it vectorized -- straight from the PTX\n", + "\n", + "We claimed the `V = 4` slice becomes one 128-bit transaction. `cute.compile[cute.KeepPTX]` keeps the\n", + "generated PTX on the result's `.artifacts.PTX`, so we read it back and look at just the global loads\n", + "and stores. Compiling the *same* kernel **with** and **without** the alignment promise shows the\n", + "difference directly: three vector instructions versus twelve scalar ones. (The wide opcode is\n", + "GPU-specific -- `v4.b32` on Ampere/Hopper, `v2.b64` on Blackwell -- but both are one 128-bit\n", + "transaction, so the *count* is the point.)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "11", + "metadata": {}, + "outputs": [], + "source": [ + "def global_mem_ops(compiled):\n", + " \"\"\"The global load/store instructions in a compiled kernel's PTX.\"\"\"\n", + " return [\n", + " ln.strip()\n", + " for ln in compiled.artifacts.PTX.splitlines()\n", + " if \"ld.global\" in ln or \"st.global\" in ln\n", + " ]\n", + "\n", + "\n", + "print(\"WITH the 16-byte promise (assumed_align=16):\")\n", + "for op in global_mem_ops(cute.compile[cute.KeepPTX](vector_add, a_, b_, c_, N, V)):\n", + " print(\" \", op)\n", + "\n", + "# Same kernel, same data -- but no alignment promise, so the compiler stays conservative.\n", + "plain = [cute.runtime.from_dlpack(t) for t in (a, b, c)]\n", + "print(\"\\nWITHOUT the promise:\")\n", + "for op in global_mem_ops(cute.compile[cute.KeepPTX](vector_add, *plain, N, V)):\n", + " print(\" \", op)\n", + "\n", + "# Expected output -- the exact opcode is GPU-specific (v4.b32 on Ampere/Hopper, v2.b64 on\n", + "# Blackwell); both are one 128-bit transaction, so the *count* is the point:\n", + "# WITH the 16-byte promise (assumed_align=16): -> 3 vector ops\n", + "# ld.global.{v4.b32 | v2.b64} ... <- one 128-bit load per operand (a and b)\n", + "# ld.global.{v4.b32 | v2.b64} ...\n", + "# st.global.{v4.b32 | v2.b64} ... <- one 128-bit store\n", + "# WITHOUT the promise: -> 12 scalar ops\n", + "# ld.global.b32 ... <- eight 32-bit loads (four per operand)\n", + "# st.global.b32 ... <- four 32-bit stores" + ] + }, + { + "cell_type": "markdown", + "id": "12", + "metadata": {}, + "source": [ + "## 1. Arrays inside `@cute.kernel`\n", + "\n", + "Inside a kernel you allocate scratch: a per-block **shared** tile and a per-thread **local** buffer.\n", + "This circular 3-tap row blur uses both. It reads a row from a **global** argument, stages it in\n", + "shared so each thread can reach its neighbors' columns, gathers three taps into local, and writes\n", + "the average back. It prints each array's layout first — a one-time, compile-time print." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "13", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def blur_kernel(inp: cutlass.Array, out: cutlass.Array, C: cutlass.Constexpr):\n", + " row, _, _ = cute.arch.block_idx() # one block per row\n", + " col, _, _ = cute.arch.thread_idx() # one thread per column\n", + "\n", + " # Step 1. Allocate scratch and print each array's layout (a one-time, compile-time print).\n", + " tile = cutlass.Array(inp.dtype, C, space=cutlass.AddressSpace.smem) # per-block scratch\n", + " taps = cutlass.Array(inp.dtype, 3) # per-thread scratch\n", + " print(\"arg :\", inp) # a global Array (the kernel argument)\n", + " print(\"shared :\", tile)\n", + " print(\"local :\", taps)\n", + "\n", + " # Step 2. Stage this row from global into shared; the barrier makes it block-visible.\n", + " row_in = inp.subview(row * inp.strides[0]) # subview: an offset view of this row\n", + " tile[col] = row_in[col] # global -> shared\n", + " cute.arch.barrier() # row now visible to every thread\n", + "\n", + " # Step 3. Gather three taps into local and write the average back.\n", + " taps[0] = tile[(col + C - 1) % C] # left neighbor (another thread wrote it)\n", + " taps[1] = tile[col]\n", + " taps[2] = tile[(col + 1) % C] # right neighbor\n", + " out[row, col] = (taps[0] + taps[1] + taps[2]) / 3.0" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "14", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def blur(inp: cutlass.Array, out: cutlass.Array, N: cutlass.Int32, C: cutlass.Constexpr):\n", + " blur_kernel(inp, out, C).launch(grid=(N, 1, 1), block=(C, 1, 1))\n", + "\n", + "\n", + "N, C = 1024, 256\n", + "inp = torch.randn(N, C, dtype=torch.float32, device=\"cuda\")\n", + "out = torch.zeros_like(inp)\n", + "\n", + "blur(cute.runtime.from_dlpack(inp), cute.runtime.from_dlpack(out), N, C=C)\n", + "\n", + "ref = (inp.roll(1, dims=1) + inp + inp.roll(-1, dims=1)) / 3.0 # circular 3-tap blur\n", + "torch.testing.assert_close(out.cpu(), ref.cpu(), atol=1e-5, rtol=1e-5)\n", + "print(\"PASS\")\n", + "\n", + "# Expected output (the layout prints fire once, while the kernel is staged):\n", + "# arg : Array(dtype=Float32, address_space=gmem, shape=(1024, 256), alignment=4)\n", + "# shared : Array(dtype=Float32, address_space=smem, shape=(256,), alignment=4)\n", + "# local : Array(dtype=Float32, address_space=generic, shape=(3,), alignment=4)\n", + "# PASS\n", + "#\n", + "# Only the local array reports `generic`, and that is correct -- not a mislabel: per-thread\n", + "# scratch is an `alloca`, stack memory addressed through a *generic* pointer, so one load/store\n", + "# path serves it. The dedicated spaces (gmem/smem/cmem) each carry their own address space and\n", + "# report it by name." + ] + }, + { + "cell_type": "markdown", + "id": "15", + "metadata": {}, + "source": [ + "## 2. Arrays outside `@cute.kernel`\n", + "\n", + "**global** and **constant** arrays are program-wide statics, so you declare them once at **module\n", + "scope** — outside any kernel — and use them by name, like CUDA C++ `__device__` / `__constant__`.\n", + "(Module-scope statics work as soon as you `import cutlass` — no extra setup.)\n", + "\n", + "`GAMMA` is a read-only constant lookup table (baked in with `init=`); `RUN_COUNT` is a global\n", + "counter that **persists across launches**. The kernel scales its input by `GAMMA` and bumps\n", + "`RUN_COUNT`. A static is a `_GlobalVariable` handle that becomes an `Array` when you index it, so we\n", + "print `.subview(0)` (the whole buffer) to see it as an `Array`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "16", + "metadata": {}, + "outputs": [], + "source": [ + "# Declared OUTSIDE any kernel, at module scope:\n", + "GAMMA = cutlass.Array(cutlass.Float32, 4, name=\"gamma\",\n", + " init=[1.0, 1.5, 2.0, 2.5], space=cutlass.AddressSpace.cmem)\n", + "RUN_COUNT = cutlass.Array(cutlass.Int32, 1, name=\"run_count\", space=cutlass.AddressSpace.gmem)\n", + "\n", + "\n", + "@cute.kernel\n", + "def scale_kernel(inp: cutlass.Array, out: cutlass.Array, seen: cutlass.Array):\n", + " i, _, _ = cute.arch.thread_idx()\n", + "\n", + " # Step 1. Materialize each module-scope static as an Array and print its layout.\n", + " print(\"constant:\", GAMMA.subview(0)) # materialize the static -> Array\n", + " print(\"global :\", RUN_COUNT.subview(0))\n", + "\n", + " # Step 2. Scale the input by the module-scope constant table.\n", + " out[i] = inp[i] * GAMMA[i] # read the module-scope constant\n", + "\n", + " # Step 3. Thread 0 bumps the module-scope global counter (it persists across launches).\n", + " if i == 0:\n", + " RUN_COUNT[0] = RUN_COUNT[0] + cutlass.Int32(1) # bump the module-scope global counter\n", + " seen[0] = RUN_COUNT[0]\n", + "\n", + "\n", + "@cute.jit\n", + "def scale(inp: cutlass.Array, out: cutlass.Array, seen: cutlass.Array):\n", + " scale_kernel(inp, out, seen).launch(grid=(1, 1, 1), block=(4, 1, 1))" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "17", + "metadata": {}, + "outputs": [], + "source": [ + "g_in = torch.tensor([2.0, 4.0, 6.0, 8.0], device=\"cuda\")\n", + "g_out = torch.zeros(4, device=\"cuda\")\n", + "seen = torch.zeros(1, dtype=torch.int32, device=\"cuda\")\n", + "\n", + "scale(cute.runtime.from_dlpack(g_in), cute.runtime.from_dlpack(g_out), cute.runtime.from_dlpack(seen))\n", + "\n", + "torch.testing.assert_close(g_out.cpu(), torch.tensor([2.0, 6.0, 12.0, 20.0])) # inp * GAMMA\n", + "assert seen.item() == 1 # counter went 0 -> 1\n", + "print(\"PASS\")\n", + "\n", + "# Expected output (layout prints fire once, while the kernel is staged):\n", + "# constant: Array(dtype=Float32, address_space=cmem, shape=(4,), alignment=4)\n", + "# global : Array(dtype=Int32, address_space=gmem, shape=(1,), alignment=4)\n", + "# PASS" + ] + }, + { + "cell_type": "markdown", + "id": "18", + "metadata": {}, + "source": [ + "## 5. Bonus — the same kernel as a decorator\n", + "\n", + "The kernel-plus-launch above is boilerplate: identical for *any* element-wise op. A small\n", + "`@elementwise` **decorator** captures it once, so a one-line op becomes a full vectorized kernel.\n", + "The op is ordinary Python that runs at trace time -- `x + y` here builds the very same `V`-wide add,\n", + "slice load/store and all." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "19", + "metadata": {}, + "outputs": [], + "source": [ + "def elementwise(op):\n", + " \"\"\"Lift a binary op into a vectorized kernel + launcher (one V-wide window per thread).\"\"\"\n", + "\n", + " @cute.kernel\n", + " def kern(a: cutlass.Array, b: cutlass.Array, c: cutlass.Array,\n", + " N: cutlass.Int32, V: cutlass.Constexpr):\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " bx, _, _ = cute.arch.block_idx()\n", + " bdx, _, _ = cute.arch.block_dim()\n", + " idx = (bx * bdx + tx) * V\n", + " if idx < N:\n", + " c[idx:V] = op(a[idx:V], b[idx:V]) # op runs at trace time, on the V-wide vectors\n", + "\n", + " @cute.jit\n", + " def host(a: cutlass.Array, b: cutlass.Array, c: cutlass.Array,\n", + " N: cutlass.Int32, V: cutlass.Constexpr):\n", + " threads = N // V\n", + " grid = ((threads + 255) // 256, 1, 1)\n", + " kern(a, b, c, N, V).launch(grid=grid, block=(256, 1, 1))\n", + "\n", + " return host\n", + "\n", + "\n", + "@elementwise\n", + "def vadd(x, y):\n", + " return x + y # the whole op -- the decorator supplies the kernel and the launch\n", + "\n", + "\n", + "vadd(a_, b_, c_, N, V)\n", + "torch.testing.assert_close(c.cpu(), (a + b).cpu(), atol=1e-5, rtol=1e-5)\n", + "print(\"PASS (decorator-generated kernel matches)\")\n", + "\n", + "# Expected output:\n", + "# PASS (decorator-generated kernel matches)" + ] + }, + { + "cell_type": "markdown", + "id": "20", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. **Drop the barrier.** Remove `cute.arch.barrier()` from `blur_kernel` — the blur goes wrong\n", + " and nondeterministic, because a thread may read a neighbor's column before that neighbor wrote it.\n", + "2. **Make the tile local.** Drop `space=…smem` from `tile` (and the barrier): each thread gets its\n", + " *own* `tile` with only `tile[col]` written, so neighbor reads hit uninitialized memory — the\n", + " difference between shared (block-wide) and local (per-thread).\n", + "3. **Inside or outside.** Move `GAMMA`'s declaration *inside* `scale_kernel` — it still works: a\n", + " named static is the same buffer wherever you declare it. Then launch `scale` twice and watch\n", + " `RUN_COUNT` climb to 2 — a global static persists across launches.\n", + "4. **Constant is read-only.** Store into `GAMMA` from a kernel — it raises `TypeError`, because\n", + " constant memory cannot be written from the device. (And `tmem` / `dsmem` from the `AddressSpace`\n", + " enum are not usable through `cutlass.Array`; they have dedicated Blackwell-chapter APIs.)" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/2_primitives/02_vector_concepts.ipynb b/examples/python/CuTeDSL/notebooks/2_primitives/02_vector_concepts.ipynb new file mode 100644 index 0000000000..d69174e890 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/2_primitives/02_vector_concepts.ipynb @@ -0,0 +1,251 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Vector concepts: `cutlass.Vector` registers vs memory\n", + "\n", + "So far data has lived in `cutlass.Array`s. But an `Array` is only a *handle to memory* -- a pointer\n", + "into global, shared, or local space. The arithmetic itself happens in **registers**, and that is\n", + "what `cutlass.Vector` is: an **immutable value** held in a thread's registers, with no address of\n", + "its own.\n", + "\n", + "Load and store connect the two:\n", + "\n", + "- **load** -- slicing an Array pulls elements into a Vector: `arr[base:V]` reads `V` contiguous\n", + " elements (registers <- memory).\n", + "- **store** -- assigning a Vector into a slice writes them back: `arr[base:V] = vec`\n", + " (memory <- registers).\n", + "\n", + "What you hold decides what you can do. An `Array` you index and assign in place; a `Vector` you can\n", + "only transform into a *new* `Vector`." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** the **Array (memory) vs Vector (registers)** split; how the slice `arr[base:V]`\n", + "(`V` is a *count*) loads a `cutlass.Vector`; element-wise SIMD arithmetic with **scalar broadcast**\n", + "(`vec * A + B`); the vectorized select `cutlass.vector.where`; a constant vector from\n", + "`cutlass.vector.full`; and the in-register **horizontal reduction** `vec.reduce(\"add\")` that\n", + "collapses all lanes to one scalar -- no shared memory, no shuffles.\n", + "\n", + "**Runs on:** any CUDA GPU. **Prereq:** the `01_array_concepts` notebook (same `arr[base:V]` slice)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. Array is memory, Vector is registers\n", + "\n", + "| type | what it is | operations |\n", + "|---|---|---|\n", + "| `cutlass.Array` | a handle to memory (global / shared / local) | `arr[i]` scalar load/store; `arr[i:V]` vector load -> `Vector`; `arr[i:V] = vec` vector store (registers -> memory) |\n", + "| `cutlass.Vector` | an immutable value in a thread's registers | `vec * A + B` element-wise + broadcast; `cutlass.vector.where(c, a, b)` vectorized select; `vec.reduce(\"add\")` horizontal fold -> scalar |\n", + "\n", + "A `Vector` is a **value, not storage**: there is no element assignment (`vec[i] = x` raises\n", + "`TypeError`). To \"change\" a Vector you apply an operation that produces a *new* one -- exactly like\n", + "adding to a Python `int`.\n", + "\n", + "The kernel below runs the whole round-trip. Each thread loads one `V`-wide tile into a Vector,\n", + "computes `act(A*x + B)` entirely in registers -- where `act` is picked from an `OPS` **dict** at\n", + "trace time -- stores the result, and also folds the tile to a single scalar with `reduce`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def vector_concepts_kernel(\n", + " x: cutlass.Array, y: cutlass.Array, s: cutlass.Array,\n", + " N: cutlass.Int32, A: cutlass.Constexpr, B: cutlass.Constexpr, V: cutlass.Constexpr,\n", + " act: cutlass.Constexpr = \"relu\",\n", + "):\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " bx, _, _ = cute.arch.block_idx()\n", + " bdx, _, _ = cute.arch.block_dim()\n", + " tidx = bx * bdx + tx\n", + " base = tidx * V # this thread owns the V-wide tile starting here\n", + "\n", + " if base < N:\n", + " # Step 1. Load: one V-wide vector load pulls the tile into registers.\n", + " # A cute.printf after each Vector op makes the per-stage transform visible (tile 0 only).\n", + " vec = x[base:V] # memory -> registers\n", + " if tidx == 0:\n", + " cute.printf(f\"load x[base:V] = [{vec[0]} {vec[1]} {vec[2]} {vec[3]}]\")\n", + "\n", + " # Step 2. Affine map: scalars A, B broadcast across all V lanes (SIMD).\n", + " affine = vec * A + B\n", + " if tidx == 0:\n", + " cute.printf(f\"affine vec*A + B = [{affine[0]} {affine[1]} {affine[2]} {affine[3]}]\")\n", + "\n", + " # Step 3. Activation. OPS maps a name to a vector op; `act` is a Constexpr, so\n", + " # OPS[act] picks the op at trace time and the kernel bakes in only that one.\n", + " zeros = cutlass.vector.full((V,), 0.0, dtype=cutlass.Float32)\n", + " OPS = {\n", + " \"relu\": lambda u: cutlass.vector.where(u > zeros, u, zeros),\n", + " \"square\": lambda u: u * u,\n", + " \"abs\": lambda u: cutlass.vector.where(u > zeros, u, zeros - u),\n", + " }\n", + " activated = OPS[act](affine)\n", + " if tidx == 0:\n", + " cute.printf(f\"act ({act}) = [{activated[0]} {activated[1]} {activated[2]} {activated[3]}]\")\n", + "\n", + " # Step 4. Horizontal fold: V lanes -> one scalar, in registers.\n", + " total = activated.reduce(\"add\")\n", + " if tidx == 0:\n", + " cute.printf(f\"reduce sum = {total}\")\n", + "\n", + " # Step 5. Store: one V-wide vector store writes the tile back, plus the per-tile sum.\n", + " y[base:V] = activated # registers -> memory\n", + " s[tidx] = total" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 2. Launch: one V-wide tile per thread\n", + "\n", + "The grid covers all `N // V` tiles with 128-thread blocks. `A`, `B`, and `V` are `Constexpr`, so the\n", + "coefficients fold into constants and the load / store / reduce become fixed-width instructions." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def vector_concepts(\n", + " x: cutlass.Array, y: cutlass.Array, s: cutlass.Array,\n", + " N: cutlass.Int32, A: cutlass.Constexpr, B: cutlass.Constexpr, V: cutlass.Constexpr,\n", + " act: cutlass.Constexpr = \"relu\",\n", + "):\n", + " block = (128, 1, 1)\n", + " tiles = N // V # one thread per V-wide tile\n", + " grid = ((tiles + block[0] - 1) // block[0], 1, 1)\n", + " vector_concepts_kernel(x, y, s, N, A, B, V, act).launch(grid=grid, block=block)" + ] + }, + { + "cell_type": "markdown", + "id": "7", + "metadata": {}, + "source": [ + "## 3. Run and check against PyTorch\n", + "\n", + "`cutlass.Array` parameters take PyTorch CUDA tensors directly via `cute.runtime.from_dlpack`\n", + "(zero-copy). We run the same kernel for each activation in the `OPS` dict and check both outputs\n", + "against the matching PyTorch op: the element-wise tile, and the per-tile row-sum." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "8", + "metadata": {}, + "outputs": [], + "source": [ + "N, V = 1 << 20, 4 # 1,048,576 elements, one 4-wide tile per thread\n", + "A, B = 2.0, 1.0 # y = act(A*x + B)\n", + "\n", + "x = torch.randn(N, dtype=torch.float32, device=\"cuda\")\n", + "\n", + "# The SAME kernel, specialized to each activation in OPS by name. `act` is a Constexpr,\n", + "# so each launch traces just one op out of the dict -- a compile-time dispatch table.\n", + "refs = {\"relu\": torch.relu, \"square\": lambda t: t * t, \"abs\": torch.abs}\n", + "for act in (\"relu\", \"square\", \"abs\"):\n", + " y = torch.zeros(N, dtype=torch.float32, device=\"cuda\")\n", + " s = torch.zeros(N // V, dtype=torch.float32, device=\"cuda\") # one row-sum per tile\n", + " vector_concepts(\n", + " cute.runtime.from_dlpack(x), cute.runtime.from_dlpack(y), cute.runtime.from_dlpack(s),\n", + " N, A, B, V, act,\n", + " )\n", + " ref = refs[act](A * x + B)\n", + " torch.testing.assert_close(y.cpu(), ref.cpu(), atol=1e-2, rtol=1e-2)\n", + " torch.testing.assert_close(s.cpu(), ref.reshape(N // V, V).sum(dim=-1).cpu(), atol=1e-2, rtol=1e-2)\n", + " print(f\"PASS act={act}\")\n", + "\n", + "# Expected output (tile 0 prints once per activation; lane values depend on the random input):\n", + "# load x[base:V] = [...]\n", + "# affine vec*A + B = [...]\n", + "# act (relu) = [...]\n", + "# reduce sum = ...\n", + "# PASS act=relu\n", + "# ... then act=square, act=abs ..." + ] + }, + { + "cell_type": "markdown", + "id": "9", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. **A Vector is immutable.** Add `vec[0] = cutlass.Float32(9.0)` -- it raises `TypeError`, because a\n", + " Vector is a value, not storage. You can only build a *new* Vector from it.\n", + "2. **Add an activation.** Add `\"gelu\": lambda u: ...` to the `OPS` dict and a matching reference --\n", + " the dict is a trace-time dispatch table, so a new entry is a new specialized kernel for free.\n", + "3. **Other reductions.** Swap `reduce(\"add\")` for `reduce(\"max\")` (and adjust the torch reference);\n", + " `reduce(\"mul\")` and `reduce(\"min\")` work too.\n", + "4. **Build a Vector from scalars.** Replace the load with\n", + " `vec = cutlass.Vector.from_elements((1.0, 2.0, 3.0, 4.0), cutlass.Float32)` -- a Vector need not\n", + " come from memory.\n", + "5. **Convert types.** Insert `vec = vec.to(cutlass.Float16)` before the affine map (store into an\n", + " `fp16` output and widen the tolerance). `Vector.to(...)` converts element types; `.bitcast(...)`\n", + " reinterprets the same bits as another type." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/2_primitives/03_cute_interop.ipynb b/examples/python/CuTeDSL/notebooks/2_primitives/03_cute_interop.ipynb new file mode 100644 index 0000000000..1c1f2bec3d --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/2_primitives/03_cute_interop.ipynb @@ -0,0 +1,250 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# `cute.Tensor` ↔ `cutlass.Array` interop (with `make_array_view`)\n", + "\n", + "The DSL gives you two handles onto the same GPU memory. A `cute.Tensor` is a pointer composed with a\n", + "CuTe **layout** — it is what `cute.runtime.from_dlpack` produces, and the type that the layout\n", + "algebra (`local_tile`, `local_partition`, ...) and copy atoms operate on. A `cutlass.Array` is a\n", + "pointer plus a **flat** shape / strides / dtype / address space — the type from the\n", + "`01_array_concepts` notebook, with its `a[idx:V]` \"slice = vectorized load/store\" idiom and the\n", + "`cutlass.Array(dtype, shape, space=...)` allocator.\n", + "\n", + "**Both** support element access (`t[(i, j)]`, `a[i, j]`) and both have a notion of slicing — but\n", + "the slices mean different things (a sub-tensor *view* vs. a `V`-wide *vector value*). Section 1\n", + "lays the two side by side.\n", + "\n", + "So far you have only seen the `cutlass.Array` side, because annotating a parameter `cutlass.Array`\n", + "and passing a `from_dlpack` tensor makes the host entry convert for you. This notebook pulls that\n", + "conversion into the open: `cutlass.make_array_view` takes a `cute.Tensor` and hands you a\n", + "`cutlass.Array` over the **same memory, no copy** — so a tensor you already hold (a kernel parameter,\n", + "a tile you sliced out) can use the Array idioms whenever they fit better." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** what a `cute.Tensor` (pointer ∘ CuTe layout — `from_dlpack`, layout algebra,\n", + "copy atoms) and a `cutlass.Array` (pointer + flat shape/strides — `a[i, j]`, `a[idx:V]` vector\n", + "slices, `space=` allocation) each support, where they overlap, where they differ, and how\n", + "`cutlass.make_array_view(tensor)` bridges the former to the latter with zero copy.\n", + "\n", + "**Runs on:** any CUDA GPU. **Prereq:** the `01_array_concepts` notebook (`cutlass.Array` indexing)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. Two handles onto GPU memory\n", + "\n", + "Neither type is a superset of the other. Both address memory element by element; they differ in\n", + "what their *layout* can express, what a *slice* returns, and which APIs consume them.\n", + "\n", + "| | `cute.Tensor` | `cutlass.Array` |\n", + "|---|---|---|\n", + "| what it is | pointer (iterator) ∘ CuTe **layout**: `T(c) = *(ptr + L(c))` | pointer + **flat** shape, strides, dtype, address space |\n", + "| layout can be | any CuTe layout: nested / hierarchical modes, composed (swizzled), static or dynamic | flat (non-nested) shape/strides only — `make_array_view` rejects nested layouts |\n", + "| produced by | `cute.runtime.from_dlpack`, `cute.make_tensor`, slicing or tiling another tensor | `cutlass.Array(dtype, shape, space=...)`, `cutlass.make_array_view(t)`, host-entry conversion of a tensor passed to a `cutlass.Array` parameter |\n", + "| element load / store | `t[i]` (linear index, mapped through the layout), `t[(i, j)]` (coordinate) | `a[i]` (flat index), `a[i, j]` (stride-aware) |\n", + "| what a slice is | `t[(i, None)]` / `cute.slice_` keeps whole modes and returns a **sub-tensor view** — no data moves | `a[idx:V]` returns a **`cutlass.Vector` of `V` contiguous elements** — a vectorized load; `V` is a *count*, not a stop index |\n", + "| vector load / store | `t.load()` / `t.store(...)` on the whole (static-layout) tensor as a `TensorSSA` | `a[idx:V]` / `a[idx:V] = vec`, with an optional `cutlass.align(16)` hint for wide transactions |\n", + "| layout algebra | `cute.local_tile`, `cute.local_partition`, `cute.zipped_divide`, `cute.composition`, `cute.recast_tensor` | none — only `a.subview(n)` (offset Array) and `a.data_ptr(n)` (raw pointer) |\n", + "| copy engines | `cute.copy` with copy atoms, incl. TMA (`make_tiled_tma_atom`) | the `05_tma_load` route: `cuda.create_tensor_map_tiled_from_view(a)` and `prims.cp_async_bulk_tensor_*` accept an Array directly |\n", + "| allocation | `cute.make_rmem_tensor` (registers), `SmemAllocator` (shared) | `cutlass.Array(..., space=rmem / smem / gmem / cmem)` |\n", + "\n", + "Rule of thumb: stay on `cute.Tensor` while you are *reshaping* memory (tiling, partitioning,\n", + "swizzling, feeding copy atoms); use `cutlass.Array` when you want the flat, CUDA-C-like access\n", + "of the `01_array_concepts` notebook — especially the `a[idx:V]` window idiom, whose result is a\n", + "`cutlass.Vector` you can do register arithmetic on.\n", + "\n", + "`cutlass.make_array_view(t)` reinterprets a tensor's base pointer and (flat) layout as a\n", + "`cutlass.Array` aliasing the **same** memory — no allocation, no copy, just a different handle onto\n", + "the same bytes.\n", + "\n", + "## 2. The bridge, inside the kernel\n", + "\n", + "This kernel takes its parameters as `cute.Tensor` on purpose — the layout-carrying type\n", + "`cute.runtime.from_dlpack` produces. It *could* read elements straight off the tensor with\n", + "`inp[idx]`; instead it wants the `01_array_concepts` window idiom, so it first builds a view with\n", + "`make_array_view`. From then on it is Array indexing: the slice `a[idx:V]` is a vectorized\n", + "load/store of `V` contiguous elements (a `cutlass.Vector`), which is what lets `* 2.0` operate on\n", + "the whole window at once.\n", + "\n", + "To make the aliasing concrete, thread 0 reads one element through the view both before and after\n", + "the store. The \"after\" value is `2 *` the \"before\" value — proof the view writes straight into the\n", + "tensor's memory." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def scale_kernel(\n", + " inp: cute.Tensor,\n", + " out: cute.Tensor,\n", + " N: cutlass.Int32,\n", + " V: cutlass.Constexpr,\n", + "):\n", + " # Step 1. This thread owns the V-wide element window starting at idx.\n", + " tx, _, _ = cute.arch.thread_idx()\n", + " bx, _, _ = cute.arch.block_idx()\n", + " bdx, _, _ = cute.arch.block_dim()\n", + " idx = (bx * bdx + tx) * V\n", + "\n", + " # Step 2. The bridge: alias each cute.Tensor as a cutlass.Array over the same\n", + " # memory (base pointer + flat layout, no copy). We could read inp[idx] off the\n", + " # tensor directly; the Array view is what gives us the a[idx:V] vector slice.\n", + " inp_arr = cutlass.make_array_view(inp)\n", + " out_arr = cutlass.make_array_view(out)\n", + "\n", + " if idx < N:\n", + " # Step 3. Scale the window. The slice is a vectorized load/store; the second\n", + " # number is a COUNT of V elements. Thread 0 reads one element through the view\n", + " # before and after to prove the view aliases the tensor's memory (out is 2x in).\n", + " if idx == 0:\n", + " cute.printf(f\"[thread 0] in : inp_arr[0]={inp_arr[0]}\")\n", + "\n", + " out_arr[idx:V] = inp_arr[idx:V] * 2.0\n", + "\n", + " if idx == 0:\n", + " cute.printf(f\"[thread 0] out: out_arr[0]={out_arr[0]} (= 2 * in)\")" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 3. Launch: one thread per V elements\n", + "\n", + "Identical to the `01_array_concepts` launch. We need `N / V` threads (one per window), in a\n", + "256-wide 1-D block with the grid rounded up." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def scale(\n", + " inp: cute.Tensor,\n", + " out: cute.Tensor,\n", + " N: cutlass.Int32,\n", + " V: cutlass.Constexpr,\n", + "):\n", + " block = (256, 1, 1)\n", + " threads = N // V # one thread per V-wide window\n", + " grid = ((threads + block[0] - 1) // block[0], 1, 1)\n", + " scale_kernel(inp, out, N, V).launch(grid=grid, block=block)" + ] + }, + { + "cell_type": "markdown", + "id": "7", + "metadata": {}, + "source": [ + "## 4. Run it and check against PyTorch\n", + "\n", + "Because the parameters are annotated `cute.Tensor`, we pass the layout-carrying values\n", + "`cute.runtime.from_dlpack` returns directly — no host-entry conversion — and the kernel does the\n", + "`make_array_view` bridge itself. Then we check against `2 * inp` on the host." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "8", + "metadata": {}, + "outputs": [], + "source": [ + "# from_dlpack yields cute.Tensors (full layout preserved); the kernel\n", + "# make_array_view's them into cutlass.Arrays for the a[idx:V] vector-slice idiom.\n", + "N, V = 1 << 20, 4\n", + "inp = torch.randn(N, dtype=torch.float32, device=\"cuda\")\n", + "out = torch.zeros(N, dtype=torch.float32, device=\"cuda\")\n", + "\n", + "scale(cute.runtime.from_dlpack(inp), cute.runtime.from_dlpack(out), N, V)\n", + "\n", + "# Verify against PyTorch (compare on host).\n", + "torch.testing.assert_close(out.cpu(), (inp * 2.0).cpu(), atol=1e-5, rtol=1e-5)\n", + "print(\"PASS\")\n", + "\n", + "# Expected output:\n", + "# PASS" + ] + }, + { + "cell_type": "markdown", + "id": "9", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. Re-annotate the parameters `cutlass.Array` and pass the same `from_dlpack` tensors. It still\n", + " works — the host entry runs this same bridge for you. So when is an *explicit* `make_array_view`\n", + " necessary? (Hint: when the `cute.Tensor` only comes into existence *inside* the kernel — e.g. a\n", + " tile from `cute.local_tile` or a `cute.slice_` — and you want Array-style access to it.)\n", + "2. Skip the bridge: replace `inp_arr[idx:V]` with a loop over `inp[idx + k]` / `out[idx + k]`\n", + " for `k in range(V)` — element access works on the tensor directly. Then compare the PTX\n", + " (`cute.compile[cute.KeepPTX]`, as in `01_array_concepts` §4): does the per-element loop\n", + " still produce the one wide load/store per window that the `a[idx:V]` slice does?\n", + "3. `make_array_view` is zero-copy. Confirm that writing through `out_arr` really changes `out` —\n", + " the view and the tensor share storage.\n", + "4. Print `inp_arr.shape` / `inp_arr.strides` / `inp_arr.dtype`: the view exposes the same\n", + " compile-time layout facts as the `cutlass.Array`s in the `01_array_concepts` notebook." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/2_primitives/04_tiled_gemm.ipynb b/examples/python/CuTeDSL/notebooks/2_primitives/04_tiled_gemm.ipynb new file mode 100644 index 0000000000..53f4688910 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/2_primitives/04_tiled_gemm.ipynb @@ -0,0 +1,287 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Tiled GEMM\n", + "\n", + "Matrix multiply, `C = A @ B`, is *the* workload GPUs are built for, and the perfect lens for the\n", + "single most important GPU performance idea: **reuse data in fast memory instead of re-reading it\n", + "from slow memory.**\n", + "\n", + "We write the kernel twice. First a **naive** version — one thread per output element — to pin\n", + "down the data flow and check it against PyTorch. It is correct but bandwidth-bound: every thread\n", + "reads a whole row of A and column of B straight from global memory. Then we stage tiles in\n", + "**shared memory**, so each element is fetched once per block instead of once per output element.\n", + "\n", + "That rewrite is your first use of `cutlass.Array(..., space=smem)` and CTA barriers — the\n", + "load -> sync -> compute -> sync loop at the heart of every fast GPU kernel." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** mapping a 2-D thread grid onto an output matrix; flat row-major indexing into\n", + "`cutlass.Array` kernel args; allocating *shaped* shared-memory tiles for 2-D access; the\n", + "cooperative-load pattern; and why the load -> sync -> compute -> sync loop needs a CTA barrier\n", + "(`cute.arch.barrier()`) on *both* ends.\n", + "\n", + "**Runs on:** any CUDA GPU. **Prereq:** the `01_array_concepts` notebook." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. The naive version (baseline)\n", + "\n", + "The simplest mapping: **one thread per output element** `C[i, j]`. Each thread reads its `(i, j)`\n", + "from the grid, walks the whole K dimension accumulating the dot product, and writes one result.\n", + "Get this working first — it pins down the indexing math and gives a correct reference for the\n", + "faster kernel.\n", + "\n", + "The matrices arrive as flat `cutlass.Array`s with no layout attached, so — just like in CUDA C —\n", + "you turn 2-D coordinates into a flat row-major offset by hand: `A[i,k] -> a[i*K+k]`,\n", + "`B[k,j] -> b[k*N+j]`, `C[i,j] -> c[i*N+j]`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def naive_gemm_kernel(\n", + " a: cutlass.Array,\n", + " b: cutlass.Array,\n", + " c: cutlass.Array,\n", + " M: cutlass.Int32,\n", + " N: cutlass.Int32,\n", + " K: cutlass.Int32,\n", + " TS: cutlass.Constexpr, # unused here; the shared host entry passes the tile size\n", + "):\n", + " tx, ty, _ = cute.arch.thread_idx()\n", + " bx, by, _ = cute.arch.block_idx()\n", + " bdx, bdy, _ = cute.arch.block_dim()\n", + "\n", + " # Step 1. This thread owns output element C[row, col].\n", + " row = bx * bdx + tx\n", + " col = by * bdy + ty\n", + "\n", + " if row < M and col < N:\n", + " # Step 2. Dot product over K. The cost: this re-reads a full row of A and a\n", + " # full column of B from global memory for every single output element.\n", + " acc = 0.0\n", + " for k in range(K):\n", + " acc += a[row * K + k] * b[k * N + col]\n", + " # Step 3. Write the one output element this thread owns.\n", + " c[row * N + col] = acc" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 2. The problem, and the fix: tile into shared memory\n", + "\n", + "In the naive kernel, neighbouring threads re-read the *same* rows of A and columns of B from\n", + "**global** memory over and over. Global memory is large but slow; the same bytes cross the bus\n", + "hundreds of times.\n", + "\n", + "**Shared memory** (SMEM) is a small, fast scratchpad shared by every thread in a block. The fix:\n", + "the block cooperatively copies a `TS x TS` tile of A and of B into SMEM *once*, then all threads\n", + "compute against that fast copy. We march along K one tile at a time, and for each tile:\n", + "\n", + "| step | what happens |\n", + "|---|---|\n", + "| **1. cooperative load** | each thread copies ONE element of the A-tile and the B-tile into SMEM |\n", + "| **2. barrier** | wait until the whole tile is in SMEM before anyone reads it |\n", + "| **3. compute** | every thread multiplies the two SMEM tiles into its accumulator |\n", + "| **4. barrier** | wait until everyone is done before the next load overwrites the tile |\n", + "\n", + "Each global element is now read **once per block** instead of once per output element — that is the\n", + "whole win. Note the two `cutlass.Array` flavors: the **global** matrices stay flat (indexed 1-D),\n", + "while the **SMEM tiles** are allocated *with a shape* `(TS, TS)` and indexed 2-D as `a_smem[ty, tx]`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def tiled_gemm_kernel(\n", + " a: cutlass.Array,\n", + " b: cutlass.Array,\n", + " c: cutlass.Array,\n", + " M: cutlass.Int32,\n", + " N: cutlass.Int32,\n", + " K: cutlass.Int32,\n", + " TS: cutlass.Constexpr,\n", + "):\n", + " tx, ty, _ = cute.arch.thread_idx()\n", + " bx, by, _ = cute.arch.block_idx()\n", + "\n", + " # Per-block scratchpads. Shaping them (TS, TS) lets us index 2-D below.\n", + " a_smem = cutlass.Array(cutlass.Float32, (TS, TS), space=cutlass.AddressSpace.smem)\n", + " b_smem = cutlass.Array(cutlass.Float32, (TS, TS), space=cutlass.AddressSpace.smem)\n", + "\n", + " # This thread owns one output element; (row, col) is its position in C.\n", + " row = bx * TS + ty\n", + " col = by * TS + tx\n", + "\n", + " acc = 0.0\n", + " for k_tile in range(0, K, TS):\n", + " # Step 1. Cooperative load: every thread brings in ONE element of each tile, so\n", + " # together the block stages the full TS x TS tiles of A and B into SMEM.\n", + " a_smem[ty, tx] = a[row * K + (k_tile + tx)]\n", + " b_smem[ty, tx] = b[(k_tile + ty) * N + col]\n", + " # Step 2. Don't read the tiles until every thread has finished writing them.\n", + " cute.arch.barrier()\n", + "\n", + " # Step 3. Multiply the two SMEM tiles into the accumulator (fast reads, no global traffic).\n", + " for kk in range(TS):\n", + " acc += a_smem[ty, kk] * b_smem[kk, tx]\n", + " # Step 4. Don't overwrite the tiles (next iteration's load) until everyone is done reading.\n", + " cute.arch.barrier()\n", + "\n", + " # Step 5. Write this thread's accumulated output element.\n", + " c[row * N + col] = acc" + ] + }, + { + "cell_type": "markdown", + "id": "7", + "metadata": {}, + "source": [ + "## 3. Launch and verify\n", + "\n", + "One block computes one `TS x TS` output tile, so the block shape *is* the tile `(TS, TS, 1)` and\n", + "the grid tiles all of C. For simplicity we assume M, N, K are multiples of `TS` (true here: 1024\n", + "with `TS=32`) — production kernels mask the ragged edge tiles instead. A small `make_gemm`\n", + "**factory** bakes the kernel variant and tile size into a ready-to-call launcher, so each kernel\n", + "gets its own `gemm` with `TS` already baked in. `cutlass.Array` parameters accept the PyTorch CUDA\n", + "tensors **directly** via `cute.runtime.from_dlpack`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "8", + "metadata": {}, + "outputs": [], + "source": [ + "def make_gemm(kernel, TS=32):\n", + " \"\"\"Factory: bake the kernel variant + tile size into a ready-to-call GEMM host.\n", + "\n", + " `kernel` and `TS` are captured at build time, so the returned host takes only the\n", + " matrices and the runtime `M, N, K` -- one specialized launcher per kernel variant.\n", + " \"\"\"\n", + "\n", + " @cute.jit\n", + " def gemm(a: cutlass.Array, b: cutlass.Array, c: cutlass.Array,\n", + " M: cutlass.Int32, N: cutlass.Int32, K: cutlass.Int32):\n", + " # One block per output tile. M, N, K stay runtime; TS is baked in, so the grid\n", + " # and block shapes are compile-time constants.\n", + " grid = (M // TS, N // TS, 1)\n", + " kernel(a, b, c, M, N, K, TS).launch(grid=grid, block=(TS, TS, 1))\n", + "\n", + " return gemm" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9", + "metadata": {}, + "outputs": [], + "source": [ + "M, N, K = 256, 256, 256\n", + "TS = 32 # output-tile size, and also the block dimension (see the tiled kernel)\n", + "# The grid and the K-loop both step by TS, so M, N, K must divide evenly by it —\n", + "# otherwise the remainder tiles would be silently dropped (no edge masking here).\n", + "assert M % TS == 0 and N % TS == 0 and K % TS == 0, \"M, N, K must each be a multiple of TS\"\n", + "a = torch.randn(M, K, dtype=torch.float32, device=\"cuda\")\n", + "b = torch.randn(K, N, dtype=torch.float32, device=\"cuda\")\n", + "ref = a.cpu() @ b.cpu()\n", + "\n", + "# Build one specialized launcher per kernel variant — TS baked in by the factory, the\n", + "# matrices and runtime M, N, K passed at the call.\n", + "naive = make_gemm(naive_gemm_kernel, TS)\n", + "tiled = make_gemm(tiled_gemm_kernel, TS)\n", + "\n", + "for name, run in ((\"naive\", naive), (\"tiled\", tiled)):\n", + " c = torch.zeros(M, N, dtype=torch.float32, device=\"cuda\")\n", + " run(cute.runtime.from_dlpack(a), cute.runtime.from_dlpack(b), cute.runtime.from_dlpack(c), M, N, K)\n", + " torch.testing.assert_close(c.cpu(), ref, atol=1e-3, rtol=1e-3)\n", + " print(f\"PASS {name}\")\n", + "\n", + "# Expected output:\n", + "# PASS naive\n", + "# PASS tiled" + ] + }, + { + "cell_type": "markdown", + "id": "10", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. Set `TS = 16` and rerun — still correct? How do the block and grid shapes change?\n", + "2. Delete the **second** `cute.arch.barrier()` and rerun a few times — the result goes wrong\n", + " intermittently. Why? (A fast thread races ahead and overwrites `a_smem` while a slower thread is\n", + " still reading the previous tile.)\n", + "3. Count the global loads per output element for the naive vs. the tiled kernel. Where did the\n", + " factor of `TS` reduction come from?" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/2_primitives/05_tma_load.ipynb b/examples/python/CuTeDSL/notebooks/2_primitives/05_tma_load.ipynb new file mode 100644 index 0000000000..a1fe1f81be --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/2_primitives/05_tma_load.ipynb @@ -0,0 +1,320 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# TMA load + mbarrier\n", + "\n", + "The **Tensor Memory Accelerator (TMA)** is a hardware copy engine that moves an entire tile from\n", + "global memory into shared memory with one instruction. No per-thread address math, no element loop:\n", + "one elected thread fires the copy, and the engine streams the bytes in the background while the rest\n", + "of the kernel runs.\n", + "\n", + "This is the smallest complete TMA program. The host builds a *descriptor*, one thread fires one\n", + "async bulk copy, every thread waits on an **mbarrier**, and we copy the loaded tile back to global\n", + "memory to check it against PyTorch. That handshake — descriptor, copy, mbarrier wait — is the\n", + "building block under every high-performance GEMM and attention kernel later in this course." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** how to build a TMA descriptor (`TensorMap`) on the host and pass it as a\n", + "`cutlass.GridConstant`; the full mbarrier handshake for one async copy\n", + "(`init → fence → barrier → arrive_expect_tx → cp_async → try_wait_parity`); why an mbarrier counts\n", + "*bytes*, not threads; and the `elect_sync` single-lane pattern for issuing the copy.\n", + "\n", + "**Runs on:** Hopper (sm_90) or newer — TMA and bulk-tensor copies are sm_90+ only. No GPU? The\n", + "kernel still compiles under `CUTE_DSL_DRYRUN=1`; the numeric check needs real hardware.\n", + "**Prereq:** the `01_array_concepts` and `02_vector_concepts` notebooks." + ] + }, + { + "cell_type": "markdown", + "id": "2", + "metadata": {}, + "source": [ + "## 1. Why TMA, and what is a descriptor?\n", + "\n", + "In the `01_array_concepts` notebook every thread computed its own global address and loaded one\n", + "element into shared memory. TMA collapses that into **one instruction**: you hand the engine a\n", + "small struct describing the global tensor and the tile to copy, and it does the rest. Two new\n", + "objects appear:\n", + "\n", + "- **TMA descriptor (`TensorMap`)** — a host-built struct describing the global tensor (shape,\n", + " strides, the tile `box` size, and the SMEM swizzle). The engine reads it to find the bytes, so\n", + " your kernel never computes an address. It is passed as `cutlass.GridConstant[cuda.TensorMap]`, so\n", + " it lives in constant memory — one copy for the whole grid.\n", + "- **mbarrier** — an `Int64` in shared memory that counts *bytes still in flight*; a copy \"arrives\"\n", + " once its bytes land. It is **not** a `cute.arch.barrier`, which synchronizes *threads*.\n", + "\n", + "The handshake for one async copy reads top to bottom:\n", + "\n", + "```text\n", + " mbarrier_init(mbar, 1) one elected thread sets the mbarrier up\n", + " fence_mbarrier_init(); barrier make the init visible to all threads\n", + " mbarrier_arrive_expect_tx(bytes) tell the mbarrier how many bytes are coming\n", + " cp_async_bulk_tensor(...) fire the TMA copy (global -> SMEM)\n", + " while not try_wait_parity(...) every thread spins until the bytes have landed\n", + "```\n", + "\n", + "First the imports. We build the descriptor with `cutlass.experimental.cuda` (imported as `cuda`),\n", + "annotate the kernel's descriptor argument with `cutlass.GridConstant`, and pull the mbarrier /\n", + "bulk-copy intrinsics from `cutlass.experimental.primitives` (imported as `prims`). The tile\n", + "size — a 64x64 fp16 tile is 8 KiB of SMEM — is set later, in the run cell." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "3", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "from cutlass.experimental import primitives as prims # mbarrier + TMA bulk-copy + barrier intrinsics\n", + "import cutlass.cute as cute\n", + "import cutlass.experimental.cuda as cuda\n", + "\n", + "import torch\n", + "import math" + ] + }, + { + "cell_type": "markdown", + "id": "4", + "metadata": {}, + "source": [ + "## 2. The kernel\n", + "\n", + "The kernel takes the descriptor and a global output tensor. It allocates the SMEM tile and a\n", + "one-entry mbarrier, runs the handshake, fires one TMA copy of the tile at global coordinate\n", + "`(0, 0)`, waits for it to land, and copies the loaded tile back to global memory.\n", + "\n", + "Three details to note:\n", + "\n", + "- **`prims.elect_sync()` picks exactly one lane.** A single TMA copy needs only one issuing thread;\n", + " the engine moves the whole tile regardless.\n", + "- **`arrive_expect_tx` declares the exact byte count** (`box elements × 2 bytes` for fp16). Without\n", + " it the wait never reaches the right count and the kernel hangs.\n", + "- **The two synchronizations do different jobs, and you need both.** `cute.arch.barrier()` waits on\n", + " *threads* — it makes the mbarrier init visible to every thread before anyone uses it. The\n", + " `try_wait_parity` spin waits on the *mbarrier* — for the copy's bytes to fully land.\n", + "\n", + "The final write-back is a plain cooperative copy — 32 threads each own two columns of the 64-wide\n", + "tile across all 64 rows — and exists only so the host can verify the load." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def tma_load_kernel(\n", + " tma_desc: cutlass.GridConstant[cuda.TensorMap],\n", + " out_tensor: cutlass.Array,\n", + " box_dim: cutlass.Constexpr,\n", + " TILE: cutlass.Constexpr = 64,\n", + "):\n", + " lane, _, _ = cute.arch.thread_idx()\n", + "\n", + " # Step 1. Allocate the SMEM destination tile and a 1-entry mbarrier for the copy.\n", + " smem_tile = cutlass.Array(cutlass.Float16, (TILE, TILE), space=cutlass.AddressSpace.smem)\n", + " mbar = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem)\n", + "\n", + " # Step 2. One lane prefetches the descriptor and arms the mbarrier for a single arrival.\n", + " if prims.elect_sync():\n", + " prims.prefetch_tensormap(tma_desc.get_ptr())\n", + " prims.mbarrier_init(mbar, 1)\n", + "\n", + " # Step 3. Publish the mbarrier init to every thread before anyone waits on it.\n", + " prims.fence_mbarrier_init()\n", + " cute.arch.barrier()\n", + "\n", + " # Step 4. One lane declares the incoming byte count, then fires the bulk copy of the\n", + " # tile at global coordinate (0, 0). arrive_expect_tx wants a byte count, so shift\n", + " # bits to bytes with `// 8`.\n", + " if prims.elect_sync():\n", + " tile_bytes = math.prod(box_dim) * cutlass.Float16.width // 8\n", + " prims.mbarrier_arrive_expect_tx(mbar, tile_bytes)\n", + " prims.cp_async_bulk_tensor_shared_cta_global(\n", + " smem_tile, tma_desc.get_ptr(), (0, 0), mbar\n", + " )\n", + "\n", + " # Step 5. Every thread spins until the tile has fully landed. The timelimit variant\n", + " # retries on a tick-timeout and uses `.acquire.cta` ordering, so the TMA writes are\n", + " # visible to all threads once the wait returns.\n", + " while not prims.mbarrier_try_wait_parity(mbar, 0):\n", + " pass\n", + "\n", + " # Step 6. Cooperative write-back SMEM -> global: lane owns columns `lane` and `lane + 32`\n", + " # across all 64 rows. Verification only.\n", + " for r in range(TILE):\n", + " out_tensor[r, lane] = smem_tile[r, lane]\n", + " out_tensor[r, lane + 32] = smem_tile[r, lane + 32]" + ] + }, + { + "cell_type": "markdown", + "id": "6", + "metadata": {}, + "source": [ + "## 3. Host: build the descriptor and launch\n", + "\n", + "`create_tensor_map_tiled_from_view` is duck-typed: it reads the source's\n", + "`shape` / `strides` / `dtype` / `data_ptr()` to program the engine, so a flat `cutlass.Array` (what\n", + "`from_dlpack` produces) feeds it directly — no separate layout object. Three arguments shape the copy:\n", + "\n", + "| argument | effect |\n", + "|---|---|\n", + "| `box_dims=box_dim[::-1]` | tile size in TMA (descriptor) order — the *reverse* of logical order |\n", + "| `stride_order=(1, 0)` | marks the innermost (contiguous) logical dimension, so kernel coordinates and the descriptor agree on the contiguous axis |\n", + "| `swizzle=cuda.TensorMapSwizzle.none` | lays the SMEM tile out exactly like the logical tile, so `smem_tile[r, c] == matrix[r, c]` and the readback compares directly against PyTorch |\n", + "\n", + "(`s128b` swizzle permutes the SMEM layout for bank-conflict-free tensor-core access — an\n", + "optimization we defer to the GEMM chapters, where the consumer is swizzle-aware.)" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "7", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def tma_load(matrix: cutlass.Array, out_tensor: cutlass.Array, TILE: cutlass.Constexpr = 64):\n", + " # Build the descriptor straight off `matrix` (shape/strides/dtype/data_ptr).\n", + " box_dim = (TILE, TILE)\n", + " tma_desc = cuda.create_tensor_map_tiled_from_view(\n", + " matrix,\n", + " box_dims=box_dim[::-1],\n", + " stride_order=(1, 0),\n", + " swizzle=cuda.TensorMapSwizzle.none,\n", + " )\n", + "\n", + " # One CTA of 32 threads is enough: a single lane issues the copy.\n", + " tma_load_kernel(tma_desc, out_tensor, box_dim, TILE).launch(\n", + " grid=(1, 1, 1), block=(32, 1, 1)\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "8", + "metadata": {}, + "source": [ + "## 4. Run it and verify\n", + "\n", + "We build a 128x128 source whose value at `[r, c]` is `r*100 + c`, TMA the top-left 64x64 tile into\n", + "SMEM, copy it back to a global output, and compare against the same slice in PyTorch.\n", + "\n", + "With `swizzle=none` and `stride_order=(1, 0)`, the loaded tile is the top-left 64x64 block in\n", + "**natural (non-transposed) order**: `got[r, c] == src[r, c]`. Row 0 reads `0, 1, 2, ...` and\n", + "column 0 reads `0, 100, 200, ...`.\n", + "\n", + "One caveat: pull the result to the host with `.cpu()` *before* comparing, and compute the\n", + "reference host-side so all comparison math runs on CPU tensors. Boolean indexing and\n", + "`assert_close` are most predictable on CPU tensors." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9", + "metadata": {}, + "outputs": [], + "source": [ + "# Source value at [r, c] is r*100 + c, so each row and column is easy to eyeball.\n", + "# Keep a CPU copy for the reference and allocate a 64x64 output for the loaded tile.\n", + "rows, cols = 128, 128\n", + "src_cpu = torch.arange(cols).unsqueeze(0) + (torch.arange(rows) * 100.0).unsqueeze(1)\n", + "src_cpu = src_cpu.to(torch.float16)\n", + "src = src_cpu.cuda()\n", + "TILE = 64\n", + "out = torch.zeros(TILE, TILE, dtype=torch.float16, device=\"cuda\")\n", + "\n", + "torch.set_printoptions(precision=1, sci_mode=False, linewidth=120)\n", + "\n", + "# Pass both torch tensors through from_dlpack as cutlass.Array. The descriptor is\n", + "# built from the static 128x128 layout, exactly what a fixed-size tutorial wants:\n", + "# the compiler bakes the concrete shape and strides into the TMA descriptor.\n", + "tma_load(cute.runtime.from_dlpack(src), cute.runtime.from_dlpack(out), TILE)\n", + "\n", + "# Verify on the host: pull to CPU first, then compare CPU tensors (see note above).\n", + "# The loaded tile is the top-left 64x64 block of `src`, in natural order.\n", + "got = out.cpu()\n", + "ref = src_cpu[:TILE, :TILE].contiguous()\n", + "print(\"loaded tile, row 0:\", got[0, :8].tolist())\n", + "print(\"loaded tile, col 0:\", got[:8, 0].tolist())\n", + "assert bool(torch.equal(got, ref)), \"TMA-loaded tile does not match the reference\"\n", + "\n", + "print(\"PASS\")\n", + "\n", + "# Expected output (value at [r, c] is r*100 + c, loaded in natural order):\n", + "# loaded tile, row 0: [0.0, 1.0, 2.0, 3.0, 4.0, 5.0, 6.0, 7.0]\n", + "# loaded tile, col 0: [0.0, 100.0, 200.0, 300.0, 400.0, 500.0, 600.0, 700.0]\n", + "# PASS" + ] + }, + { + "cell_type": "markdown", + "id": "10", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. Change the copy coordinate from `(0, 0)` to `(0, 64)` in the kernel — which 64x64 tile of `src`\n", + " lands in SMEM now, and how should the reference slice change to match?\n", + "2. Delete the `mbarrier_arrive_expect_tx` line. The wait can no longer reach the right byte count —\n", + " what happens, and why does every TMA copy pair with an expect-tx?\n", + "3. The mbarrier (an `Int64` in SMEM) tracks *bytes in flight*, while `cute.arch.barrier` syncs\n", + " *threads*. Explain in one sentence why this kernel needs both.\n", + "4. Switch `swizzle` to `cuda.TensorMapSwizzle.s128b` and rerun. The numeric check fails because the\n", + " SMEM bytes are now permuted — a preview of why GEMM kernels keep a swizzle-aware consumer.\n", + "\n", + "## Peeking under the hood\n", + "\n", + "Set a `CUTE_DSL_*` option before running the launch cell — in a fresh cell run\n", + "`import os; os.environ[\"CUTE_DSL_PRINT_PTX\"] = \"1\"`, then re-execute the run cell — to see the\n", + "IR / PTX the DSL generates:\n", + "\n", + "- `CUTE_DSL_DRYRUN=1` — trace and compile only, no GPU\n", + "- `CUTE_DSL_PRINT_IR=1` — the generated IR (look for the `cp.async.bulk` op)\n", + "- `CUTE_DSL_PRINT_PTX=1` — PTX (`cp.async.bulk.tensor` + mbarrier ops)" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/2_primitives/06_prefix_sum.ipynb b/examples/python/CuTeDSL/notebooks/2_primitives/06_prefix_sum.ipynb new file mode 100644 index 0000000000..fb046dd234 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/2_primitives/06_prefix_sum.ipynb @@ -0,0 +1,248 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Prefix sum (warp-shuffle scan)\n", + "\n", + "Our first **non-deep-learning** primitive. A *scan* turns a row `[a, b, c, d]` into its running\n", + "totals `[a, a+b, a+b+c, a+b+c+d]`. It is the quiet engine behind sorting, stream compaction, and\n", + "histograms.\n", + "\n", + "Each output depends on every element before it, so a scan looks stubbornly *sequential*. The key\n", + "insight of this notebook is that it isn't: the very same **two-level (warp → shared memory)**\n", + "reduction you built for softmax also parallelizes a scan. The only twist is that we keep *every*\n", + "partial sum along the way instead of collapsing the row down to a single number." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** the inclusive parallel scan; the **Hillis-Steele warp scan** built from\n", + "`cute.arch.shuffle_sync_up` in log₂(32) = 5 steps -- factored into a small reusable `@cute.jit`\n", + "helper; and how to stitch per-warp scans into a block-wide scan with one shared-memory pass.\n", + "\n", + "**Runs on:** any CUDA GPU. **Prereq:** the `01_softmax` notebook (same two-level reduction)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. A scan in two levels\n", + "\n", + "One block per row, one thread per element, so `BLOCK_SIZE == C`. The strategy mirrors the softmax\n", + "reduction, but a scan must keep *every* intermediate sum instead of collapsing the row to one\n", + "number.\n", + "\n", + "```text\n", + " 1. WARP SCAN each warp inclusive-scans its 32 lanes with shuffle_sync_up (offsets 1,2,4,8,16)\n", + " 2. WARP TOTALS the last lane of each warp writes its running total to shared memory\n", + " 3. SCAN TOTALS warp 0 scans those per-warp totals (the very same warp scan)\n", + " 4. ADD PREFIX each warp adds the totals of all earlier warps -> the full row scan\n", + "```\n", + "\n", + "The engine of step 1 is `shuffle_sync_up(v, d)`, which hands a lane the value held `d` lanes below\n", + "it. A lane folds that value in only when `lane >= d`; otherwise there is no real neighbor that far\n", + "down. Doubling the distance reaches every lane below in log₂(32) = 5 steps:\n", + "\n", + "| step | offset | a lane now holds |\n", + "|---|---|---|\n", + "| 1 | 1 | sum of itself + 1 lane below |\n", + "| 2 | 2 | sum of the 4 nearest lanes |\n", + "| 3 | 4 | sum of the 8 nearest lanes |\n", + "| 4 | 8 | sum of the 16 nearest lanes |\n", + "| 5 | 16 | sum of all 32 lanes at and below it |\n", + "\n", + "Steps 1 and 3 run the **identical** warp scan, so we write it once as a `@cute.jit` helper and call\n", + "it at both levels -- a reusable *device subroutine*, the kind of factoring you'll lean on in the\n", + "Blackwell chapters." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def warp_inclusive_scan(val, lane_id):\n", + " \"\"\"Inclusive prefix sum up the 32 lanes of a warp (Hillis-Steele, 5 shuffle steps).\"\"\"\n", + " for offset in [1, 2, 4, 8, 16]:\n", + " below = cute.arch.shuffle_sync_up(val, offset) # value held `offset` lanes below\n", + " if lane_id >= offset: # fold in only a real neighbor\n", + " val = val + below\n", + " return val" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def prefix_sum_kernel(\n", + " inp: cutlass.Array, out: cutlass.Array, C: cutlass.Constexpr, BLOCK_SIZE: cutlass.Constexpr\n", + "):\n", + " WARP_SIZE = 32\n", + " WARPS_PER_BLOCK = BLOCK_SIZE // WARP_SIZE\n", + " warp_totals = cutlass.Array(cutlass.Float32, WARPS_PER_BLOCK, space=cutlass.AddressSpace.smem)\n", + "\n", + " row, _, _ = cute.arch.block_idx() # one block per row\n", + " tid, _, _ = cute.arch.thread_idx()\n", + " warp_id = tid // WARP_SIZE\n", + " lane_id = tid % WARP_SIZE\n", + " base = row * C # flat (N, C): one element per thread (C == BLOCK_SIZE)\n", + "\n", + " # Step 1. Inclusive scan within the warp (the helper).\n", + " running_sum = warp_inclusive_scan(inp[base + tid], lane_id)\n", + "\n", + " # Step 2. The last lane holds the warp total; publish it for the cross-warp scan.\n", + " if lane_id == WARP_SIZE - 1:\n", + " warp_totals[warp_id] = running_sum\n", + " cute.arch.barrier()\n", + "\n", + " # Step 3. Warp 0 scans the per-warp totals with the SAME helper, so warp_totals[w]\n", + " # becomes the sum of warps 0..w. Idle lanes carry 0.\n", + " if warp_id == 0:\n", + " warp_total = 0.0\n", + " if lane_id < WARPS_PER_BLOCK:\n", + " warp_total = warp_totals[lane_id]\n", + " warp_total = warp_inclusive_scan(warp_total, lane_id)\n", + " if lane_id < WARPS_PER_BLOCK:\n", + " warp_totals[lane_id] = warp_total\n", + " cute.arch.barrier()\n", + "\n", + " # Step 4. Lift each warp-local scan to the row scan by adding earlier warps' totals.\n", + " if warp_id > 0:\n", + " running_sum = running_sum + warp_totals[warp_id - 1]\n", + " out[base + tid] = running_sum" + ] + }, + { + "cell_type": "markdown", + "id": "6", + "metadata": {}, + "source": [ + "## 2. Launch: one block per row\n", + "\n", + "Grid `(N, 1, 1)` is one block per row; block `(C, 1, 1)` is one thread per column, so\n", + "`BLOCK_SIZE == C`. Because `C` is a `Constexpr`, the shuffle offsets and the warp count are known at\n", + "compile time, so the scan loops unroll into straight-line shuffle code." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "7", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def prefix_sum(\n", + " inp: cutlass.Array,\n", + " out: cutlass.Array,\n", + " N: cutlass.Int32,\n", + " C: cutlass.Constexpr,\n", + "):\n", + " # One block per row, one thread per column (BLOCK_SIZE == C).\n", + " prefix_sum_kernel(inp, out, C, C).launch(grid=(N, 1, 1), block=(C, 1, 1))" + ] + }, + { + "cell_type": "markdown", + "id": "8", + "metadata": {}, + "source": [ + "## 3. Run and check against PyTorch\n", + "\n", + "`cutlass.Array` parameters accept the PyTorch CUDA tensors **directly** through\n", + "`cute.runtime.from_dlpack` -- no manual copy. The reference is `torch.cumsum`, which is exactly an\n", + "inclusive scan with addition." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9", + "metadata": {}, + "outputs": [], + "source": [ + "# Inputs live on the GPU and are handed to the kernel straight through DLPack.\n", + "N, C = 1024, 256\n", + "inp = torch.randn(N, C, dtype=torch.float32, device=\"cuda\")\n", + "out = torch.zeros_like(inp)\n", + "\n", + "prefix_sum(cute.runtime.from_dlpack(inp), cute.runtime.from_dlpack(out), N, C=C)\n", + "\n", + "# torch.cumsum is the inclusive prefix sum; compare on the host.\n", + "torch.testing.assert_close(\n", + " out.cpu(), torch.cumsum(inp.cpu(), dim=-1), atol=1e-2, rtol=1e-3\n", + ")\n", + "print(\"PASS\")" + ] + }, + { + "cell_type": "markdown", + "id": "10", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. **Exclusive scan** (`out[i] = sum(inp[0..i-1])`, with `out[0] = 0`): fold the neighbor in only\n", + " when `lane > offset`, or shift the inclusive result right by one lane. Check it against a shifted\n", + " `torch.cumsum`.\n", + "2. **A different operator:** swap `+` for `max` and the scan becomes a running maximum — the same\n", + " structure with a different monoid. What identity value replaces the `0.0` seed in warp 0?\n", + "3. **Rows wider than a block (`C > BLOCK_SIZE`):** scan the row in tiles of `BLOCK_SIZE`, carrying\n", + " each tile's grand total into the next. Which part is forced to be sequential, and which stays\n", + " parallel?" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/3_kernels/01_softmax.ipynb b/examples/python/CuTeDSL/notebooks/3_kernels/01_softmax.ipynb new file mode 100644 index 0000000000..9eb694b710 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/3_kernels/01_softmax.ipynb @@ -0,0 +1,465 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Softmax (naive → online/flash)\n", + "\n", + "Softmax is the first kernel in this course where *how* you reduce matters as much as *what* you\n", + "compute. A correct GPU softmax has to dodge two traps at once: floating-point **overflow**, and\n", + "**wasted passes** over the row.\n", + "\n", + "We build it twice. First a straightforward **naive** kernel — three passes over the row, with a\n", + "classic shared-memory **tree reduction** for the max and the sum. It is easy to read and a good\n", + "mental model, but it touches the data more than it needs to. Then we build the **online (flash)**\n", + "kernel that fuses the max and the sum into a *single* pass, and reduces across the block with\n", + "warp shuffles plus a tiny shared-memory hand-off — the same two-level reduction pattern you will\n", + "reuse in later notebooks." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** why numerically-stable softmax subtracts the row max; a shared-memory **tree\n", + "reduction** (the naive baseline); the online `(max, sum)` recurrence and why merging two partial\n", + "states *rescales* one onto the other -- captured once in a `merge()` helper; and a **two-level\n", + "reduction** (warp shuffle then shared memory).\n", + "\n", + "**Runs on:** any CUDA GPU. **Prereq:** the `01_array_concepts` and `02_vector_concepts` notebooks." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. Stable softmax needs the row max\n", + "\n", + "The definition is `softmax(x)_i = exp(x_i) / sum_j exp(x_j)`. Computed literally, `exp(x_i)`\n", + "overflows float32 the moment any `x_i` is even moderately large — `exp(89)` is already past the\n", + "float32 ceiling. The fix is the **numerically-stable** form, which subtracts the row max\n", + "`m = max_j x_j` before exponentiating:\n", + "\n", + "```text\n", + "softmax(x)_i = exp(x_i - m) / sum_j exp(x_j - m)\n", + "```\n", + "\n", + "This is the *same* math — the `exp(-m)` factor cancels between numerator and denominator — but\n", + "every exponent is now `<= 0`, so each `exp(...)` lands in `(0, 1]` and can never overflow. The\n", + "price is that, done naively, it reads the row **three times**: once to find `m`, once to sum\n", + "`exp(x - m)`, and once to normalize." + ] + }, + { + "cell_type": "markdown", + "id": "4", + "metadata": {}, + "source": [ + "## 2. Two ways to reduce the row\n", + "\n", + "| | naive (this section) | online / flash (section 5) |\n", + "|---|---|---|\n", + "| passes over the row | **3** — max, then exp, then sum | **1** — `(max, sum)` fused |\n", + "| reduction | shared-memory **tree** | warp **shuffle** + a small shared hand-off |\n", + "| extra global traffic | the `exp`s round-trip through global | none |\n", + "\n", + "The naive kernel does exactly what the math says, in three passes; a shared-memory tree reduction\n", + "collapses the per-thread maxes (then sums) to one value. It's easy to read but touches the row three\n", + "times. Section 5 fuses all three into a single pass." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "5", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def softmax_naive(inp: cutlass.Array, out: cutlass.Array, C: cutlass.Constexpr, BLOCK_SIZE: cutlass.Constexpr):\n", + " \"\"\"Three-pass softmax over one row, using shared-memory tree reductions.\"\"\"\n", + " shared = cutlass.Array(cutlass.Float32, BLOCK_SIZE, space=cutlass.AddressSpace.smem)\n", + " row, _, _ = cute.arch.block_idx() # one block per row\n", + " tid, _, _ = cute.arch.thread_idx()\n", + " base = row * C # flat (N, C) addressing: row * C + col\n", + "\n", + " # Step 1. Each thread maxes over its strided slice of the row.\n", + " maxval = -3.4028235e38 # -FLT_MAX, the identity for max\n", + " for i in cutlass.range_constexpr(0, C, BLOCK_SIZE):\n", + " col = i + tid\n", + " if col < C:\n", + " maxval = cute.math.max(maxval, inp[base + col])\n", + "\n", + " # Step 2. Tree-reduce the per-thread maxes down to shared[0].\n", + " shared[tid] = maxval\n", + " for stride in [64, 32, 16, 8, 4, 2, 1]:\n", + " cute.arch.barrier()\n", + " if stride < BLOCK_SIZE and tid < stride:\n", + " shared[tid] = cute.math.max(shared[tid], shared[tid + stride])\n", + " cute.arch.barrier()\n", + " row_max = shared[0]\n", + "\n", + " # Step 3. Write the stable exponentials exp(x - row_max).\n", + " for i in cutlass.range_constexpr(0, C, BLOCK_SIZE):\n", + " col = i + tid\n", + " if col < C:\n", + " out[base + col] = cute.math.exp(inp[base + col] - row_max, fastmath=True)\n", + " cute.arch.barrier()\n", + "\n", + " # Step 4. Each thread sums its strided slice of the exps it just wrote.\n", + " sumval = 0.0\n", + " for i in cutlass.range_constexpr(0, C, BLOCK_SIZE):\n", + " col = i + tid\n", + " if col < C:\n", + " sumval = sumval + out[base + col]\n", + "\n", + " # Step 5. Tree-reduce the per-thread sums down to shared[0].\n", + " shared[tid] = sumval\n", + " for stride in [64, 32, 16, 8, 4, 2, 1]:\n", + " cute.arch.barrier()\n", + " if stride < BLOCK_SIZE and tid < stride:\n", + " shared[tid] = shared[tid] + shared[tid + stride]\n", + " cute.arch.barrier()\n", + " row_sum = shared[0]\n", + "\n", + " # Step 6. Normalize each exp by the row sum.\n", + " for i in cutlass.range_constexpr(0, C, BLOCK_SIZE):\n", + " col = i + tid\n", + " if col < C:\n", + " out[base + col] = out[base + col] / row_sum" + ] + }, + { + "cell_type": "markdown", + "id": "6", + "metadata": {}, + "source": [ + "## 3. Launch the naive kernel\n", + "\n", + "The host entry point just launches the kernel: one block per row (`grid = (N, 1, 1)`) with\n", + "`BLOCK_SIZE` threads each. Because `C` is a `Constexpr`, the strided `range_constexpr` loops have\n", + "a compile-time trip count and are unrolled." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "7", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def softmax_naive_host(\n", + " inp: cutlass.Array,\n", + " out: cutlass.Array,\n", + " N: cutlass.Int32,\n", + " C: cutlass.Constexpr,\n", + " BLOCK_SIZE: cutlass.Constexpr,\n", + "):\n", + " softmax_naive(inp, out, C, BLOCK_SIZE).launch(\n", + " grid=(N, 1, 1), block=(BLOCK_SIZE, 1, 1)\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "8", + "metadata": {}, + "source": [ + "## 4. Run and check the naive kernel\n", + "\n", + "`cutlass.Array` kernel parameters accept PyTorch CUDA tensors **directly** through\n", + "`cute.runtime.from_dlpack` — no manual copy or pointer wrangling. We check the result against\n", + "`torch.softmax`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "9", + "metadata": {}, + "outputs": [], + "source": [ + "N, C = 1024, 2048\n", + "inp = torch.randn(N, C, dtype=torch.float32, device=\"cuda\")\n", + "out = torch.zeros_like(inp)\n", + "\n", + "softmax_naive_host(\n", + " cute.runtime.from_dlpack(inp),\n", + " cute.runtime.from_dlpack(out),\n", + " N,\n", + " C=C,\n", + " BLOCK_SIZE=128,\n", + ")\n", + "\n", + "torch.testing.assert_close(\n", + " out.cpu(), torch.softmax(inp.cpu(), dim=-1), atol=1e-3, rtol=1e-3\n", + ")\n", + "print(\"PASS Naive\")" + ] + }, + { + "cell_type": "markdown", + "id": "10", + "metadata": {}, + "source": [ + "## 5. The online (flash) trick: one pass for (max, sum)\n", + "\n", + "The naive kernel needs a separate pass for the max because the stable sum `sum exp(x - m)`\n", + "depends on `m`. The **online** algorithm breaks that dependency: it computes the max and the sum\n", + "*together* in one pass by carrying a running pair `(m, s)` and **rescaling** the sum whenever the\n", + "max grows (Milakov & Gimelshein, *Online normalizer calculation for softmax*, arXiv:1805.02867).\n", + "For each new value `x`:\n", + "\n", + "```text\n", + "m_new = max(m, x)\n", + "s_new = s * exp(m - m_new) + exp(x - m_new)\n", + "```\n", + "\n", + "The `s * exp(m - m_new)` factor is the whole idea: every term already in `s` was measured against\n", + "the *old* max, so when the max moves we rescale the accumulated sum onto the new reference.\n", + "\n", + "The same rescale lets us **merge two partial states** `(m_a, s_a)` and `(m_b, s_b)` that were\n", + "computed over disjoint chunks of the row — and it is why the merge is a rescale, **never a bare\n", + "add**:\n", + "\n", + "```text\n", + "m = max(m_a, m_b)\n", + "s = s_a * exp(m_a - m) + s_b * exp(m_b - m)\n", + "```\n", + "\n", + "**This kernel: one block per row, reduced in two levels.** Each thread folds its *strided* slice\n", + "of the row into a private `(m, s)` (one online pass). Threads then merge within a warp using\n", + "`shuffle_sync_down` (offsets 16, 8, 4, 2, 1 — no shared memory, no barrier), and finally across\n", + "warps through a small shared-memory array guarded by a CTA barrier. Once the block agrees on the\n", + "row's `(m, s)`, every thread writes `exp(x - m) / s`.\n", + "\n", + "```text\n", + " ONE BLOCK PER ROW each thread t: online-fold a strided slice -> (m_t, s_t)\n", + " STEP 1 warp reduce via shuffle_sync_down, offsets 16,8,4,2,1 (no smem, no barrier)\n", + " warp0 -> lane0:(M0,S0) warp1 -> lane0:(M1,S1) ... merge = max + exp-rescale\n", + " STEP 2 smem[warp_id] = (M_w, S_w); barrier; thread0 merges all warps -> (M,S); barrier\n", + " FINAL every thread writes out_i = exp(x_i - M) / S\n", + "```" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "11", + "metadata": {}, + "outputs": [], + "source": [ + "def merge(m, s, other_m, other_s):\n", + " \"\"\"Merge two online-softmax (max, sum) states onto a common max -- never a bare add.\n", + "\n", + " new_max = max(m, other_m), then BOTH sums are rescaled onto it:\n", + " new_sum = s*exp(m - new_max) + other_s*exp(other_m - new_max)\n", + " The exp(... - new_max) factors are the rescale. Folding a single value x is just the\n", + " same merge against the one-element state (x, 1): merge(m, s, x, 1.0).\n", + " \"\"\"\n", + " new_max = cute.math.max(m, other_m)\n", + " new_sum = (s * cute.math.exp(m - new_max, fastmath=True)\n", + " + other_s * cute.math.exp(other_m - new_max, fastmath=True))\n", + " return new_max, new_sum" + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "12", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def softmax_kernel(inp: cutlass.Array, out: cutlass.Array, C: cutlass.Constexpr, BLOCK_SIZE: cutlass.Constexpr):\n", + " \"\"\"Row-wise numerically-stable softmax in one online (max, sum) pass.\n", + "\n", + " Inputs:\n", + " inp -- Float32 array viewing a flat (N, C) row-major buffer; this block reads row\n", + " block_idx.x, the C contiguous elements starting at base = row * C.\n", + " out -- Float32 array over the same (N, C) buffer; written in place (same shape/dtype).\n", + " C -- row width (Constexpr), so the strided range_constexpr loops unroll.\n", + " BLOCK_SIZE -- threads per block (Constexpr); a multiple of the 32-lane warp size.\n", + "\n", + " Output: writes this block's row of `out` with softmax(x)_i = exp(x_i - row_max) / row_sum,\n", + " where row_max = max_j x_j and row_sum = sum_j exp(x_j - row_max). Each thread folds its\n", + " strided slice into a private (max, sum) via merge(); the block reduces those partials in\n", + " two levels (warp shuffle, then a shared-memory hand-off), and every thread normalizes.\n", + " \"\"\"\n", + " WARP_SIZE = 32\n", + " WARPS_PER_BLOCK = BLOCK_SIZE // WARP_SIZE\n", + " NEG_INF = -3.4028235e38 # -FLT_MAX, the running-max identity\n", + "\n", + " # One (max, sum) partial per warp, handed off through shared memory.\n", + " smax = cutlass.Array(cutlass.Float32, WARPS_PER_BLOCK, space=cutlass.AddressSpace.smem)\n", + " ssum = cutlass.Array(cutlass.Float32, WARPS_PER_BLOCK, space=cutlass.AddressSpace.smem)\n", + "\n", + " row, _, _ = cute.arch.block_idx() # one block per row\n", + " tid, _, _ = cute.arch.thread_idx()\n", + " warp_id = tid // WARP_SIZE\n", + " lane_id = tid % WARP_SIZE\n", + " base = row * C\n", + "\n", + " # Step 1. Fold this thread's strided slice into a running (max, sum), one value at a time.\n", + " m, s = NEG_INF, 0.0\n", + " for i in cutlass.range_constexpr(0, C, BLOCK_SIZE):\n", + " col = i + tid\n", + " if col < C:\n", + " m, s = merge(m, s, inp[base + col], 1.0)\n", + "\n", + " # Step 2. Reduce across the warp with shuffles -- lane 0 ends with the warp's partial.\n", + " for offset in [16, 8, 4, 2, 1]:\n", + " m, s = merge(m, s, cute.arch.shuffle_sync_down(m, offset), cute.arch.shuffle_sync_down(s, offset))\n", + "\n", + " # Step 3. Each warp publishes its partial; the barrier makes them all visible.\n", + " if lane_id == 0:\n", + " smax[warp_id] = m\n", + " ssum[warp_id] = s\n", + " cute.arch.barrier()\n", + "\n", + " # Step 4. Thread 0 merges the per-warp partials into the row (max, sum).\n", + " if tid == 0:\n", + " m, s = smax[0], ssum[0]\n", + " for w in cutlass.range_constexpr(1, WARPS_PER_BLOCK):\n", + " m, s = merge(m, s, smax[w], ssum[w])\n", + " smax[0], ssum[0] = m, s\n", + " cute.arch.barrier()\n", + " row_max, row_sum = smax[0], ssum[0]\n", + "\n", + " # Step 5. Write the normalized softmax.\n", + " for i in cutlass.range_constexpr(0, C, BLOCK_SIZE):\n", + " col = i + tid\n", + " if col < C:\n", + " out[base + col] = cute.math.exp(inp[base + col] - row_max, fastmath=True) / row_sum" + ] + }, + { + "cell_type": "markdown", + "id": "13", + "metadata": {}, + "source": [ + "## 6. Launch: one block per row\n", + "\n", + "Same launch shape as before — grid `(N, 1, 1)` is one block per row — but with `BLOCK_SIZE = 256`\n", + "threads, i.e. 8 warps per block. As before, `C` is a `Constexpr`, so the strided\n", + "`range_constexpr` loops are unrolled." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "14", + "metadata": {}, + "outputs": [], + "source": [ + "@cute.jit\n", + "def softmax(\n", + " inp: cutlass.Array,\n", + " out: cutlass.Array,\n", + " N: cutlass.Int32,\n", + " C: cutlass.Constexpr,\n", + " BLOCK_SIZE: cutlass.Constexpr,\n", + "):\n", + " softmax_kernel(inp, out, C, BLOCK_SIZE).launch(\n", + " grid=(N, 1, 1), block=(BLOCK_SIZE, 1, 1)\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "15", + "metadata": {}, + "source": [ + "## 7. Run and check against PyTorch\n", + "\n", + "Same setup as the naive run: pass the CUDA tensors straight through `cute.runtime.from_dlpack`\n", + "and compare against `torch.softmax`. The online kernel is a single fused pass plus the two-level\n", + "reduction, yet produces the identical result." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "16", + "metadata": {}, + "outputs": [], + "source": [ + "N, C = 1024, 2048\n", + "inp = torch.randn(N, C, dtype=torch.float32, device=\"cuda\")\n", + "out = torch.zeros_like(inp)\n", + "\n", + "softmax(\n", + " cute.runtime.from_dlpack(inp),\n", + " cute.runtime.from_dlpack(out),\n", + " N,\n", + " C=C,\n", + " BLOCK_SIZE=256,\n", + ")\n", + "\n", + "torch.testing.assert_close(\n", + " out.cpu(), torch.softmax(inp.cpu(), dim=-1), atol=1e-3, rtol=1e-3\n", + ")\n", + "print(\"PASS Optimized\")" + ] + }, + { + "cell_type": "markdown", + "id": "17", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. Change `C` (e.g. to 4096) and rerun -- the `range_constexpr` trip count changes. Does it still\n", + " match torch?\n", + "2. **Why the merge rescales.** Build a 2-element example where merging two `(max, sum)` states by\n", + " *adding* the sums (no `exp(Δmax)` rescale) gives the wrong answer -- the concrete reason `merge()`\n", + " rescales instead of adding.\n", + "3. **One merge, three places.** `merge()` is called in the per-thread fold, the warp-shuffle reduce,\n", + " and the cross-warp reduce. Trace one row's `(max, sum)` through all three and confirm the result\n", + " is the row's true max and `sum(exp(x - max))`." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/3_kernels/02_stencil_2d.ipynb b/examples/python/CuTeDSL/notebooks/3_kernels/02_stencil_2d.ipynb new file mode 100644 index 0000000000..3b2eb70c3d --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/3_kernels/02_stencil_2d.ipynb @@ -0,0 +1,241 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "\n", + "\n", + "# 2D stencil: a heat-diffusion step\n", + "\n", + "A **stencil** rewrites every grid point as a weighted sum of itself and its neighbors. The\n", + "classic example is a single **Jacobi heat-diffusion step** on a 2D grid, which is also just a\n", + "5-point local average. We write it two ways against the same reference: a straightforward\n", + "global-memory kernel, then a shared-memory-tiled version that stages each block's tile **plus a\n", + "one-cell halo** into SMEM so neighbor reads stay on-chip." + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "**You'll learn:** 2-D access over a flat `cutlass.Array` (`a[row*W + col]`); reading a point's\n", + "four neighbors with **zero-boundary** handling (out-of-grid neighbors count as 0); and the\n", + "**halo** pattern — staging a `(TS+2)x(TS+2)` shared-memory tile (the block's `TS x TS` output\n", + "region ringed by one row/column of neighbors) so the stencil reads from SMEM instead of\n", + "re-reading global memory up to five times per point.\n", + "\n", + "**Runs on:** any CUDA GPU." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass\n", + "import cutlass.cute as cute\n", + "import torch\n", + "import torch.nn.functional as F\n", + "\n", + "# One Jacobi heat step == a 5-point average; out-of-grid neighbors count as 0:\n", + "# out[i,j] = C0*in[i,j] + C1*(in[i-1,j] + in[i+1,j] + in[i,j-1] + in[i,j+1])\n", + "C0 = 0.2 # center weight (baked at trace time -- plain Python floats)\n", + "C1 = 0.2 # neighbor weight\n", + "TS = 16 # block is TS x TS; H and W must be multiples of TS" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 1. The naive kernel\n", + "\n", + "One thread owns one output point `(row, col)` and reads its four neighbors straight from global\n", + "memory. The boundary is handled by **skipping** any neighbor that falls off the grid — a skipped\n", + "neighbor contributes nothing, i.e. it counts as 0. Each `if` is a plain Python `if` on a *staged*\n", + "value, so the DSL turns it into a real GPU branch and never reads out of bounds." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def stencil_naive_kernel(inp: cutlass.Array, out: cutlass.Array, H: cutlass.Int32, W: cutlass.Int32):\n", + " tx, ty, _ = cute.arch.thread_idx()\n", + " bx, by, _ = cute.arch.block_idx()\n", + " bdx, bdy, _ = cute.arch.block_dim()\n", + " row = bx * bdx + tx\n", + " col = by * bdy + ty\n", + " if row < H and col < W:\n", + " acc = C0 * inp[row * W + col]\n", + " # Add each in-grid neighbor; an out-of-grid neighbor is simply skipped (counts as 0).\n", + " if row > 0:\n", + " acc = acc + C1 * inp[(row - 1) * W + col]\n", + " if row < H - 1:\n", + " acc = acc + C1 * inp[(row + 1) * W + col]\n", + " if col > 0:\n", + " acc = acc + C1 * inp[row * W + (col - 1)]\n", + " if col < W - 1:\n", + " acc = acc + C1 * inp[row * W + (col + 1)]\n", + " out[row * W + col] = acc" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 2. Shared-memory tiling with a halo\n", + "\n", + "The naive kernel re-reads every point up to five times from global memory (each point is a\n", + "neighbor of four others). A block can instead stage its data **once** into shared memory and then\n", + "read neighbors on-chip. The catch: threads on the block's edge need neighbors that belong to the\n", + "*next* block — the **halo**. So the SMEM tile is `(TS+2) x (TS+2)`: the block's `TS x TS` output\n", + "region plus a one-cell ring around it.\n", + "\n", + "Loading is cooperative — every thread writes its own center cell, and the edge threads\n", + "(`tx == 0`, `tx == TS-1`, `ty == 0`, `ty == TS-1`) additionally fetch the halo rows/columns,\n", + "using 0 where the halo falls off the grid. A `cute.arch.barrier()` then makes the full tile\n", + "visible before anyone reads it, and the stencil becomes four SMEM lookups." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "@cute.kernel\n", + "def stencil_smem_kernel(inp: cutlass.Array, out: cutlass.Array, H: cutlass.Int32, W: cutlass.Int32):\n", + " tx, ty, _ = cute.arch.thread_idx()\n", + " bx, by, _ = cute.arch.block_idx()\n", + " # SMEM tile = the block's TS x TS region + a 1-cell halo on every side.\n", + " tile = cutlass.Array(cutlass.Float32, (TS + 2, TS + 2), space=cutlass.AddressSpace.smem)\n", + " row0 = bx * TS # this block's top-left output cell in the grid\n", + " col0 = by * TS\n", + " row = row0 + tx\n", + " col = col0 + ty\n", + "\n", + " # Center: every thread brings in its own cell (in-bounds since H, W are multiples of TS).\n", + " tile[tx + 1, ty + 1] = inp[row * W + col]\n", + " # Top / bottom halo rows: loaded by the first / last row of threads (0 if off the grid).\n", + " if tx == 0:\n", + " v = cutlass.Float32(0.0)\n", + " if row0 - 1 >= 0:\n", + " v = inp[(row0 - 1) * W + col]\n", + " tile[0, ty + 1] = v\n", + " if tx == TS - 1:\n", + " v = cutlass.Float32(0.0)\n", + " if row0 + TS < H:\n", + " v = inp[(row0 + TS) * W + col]\n", + " tile[TS + 1, ty + 1] = v\n", + " # Left / right halo columns: loaded by the first / last column of threads.\n", + " if ty == 0:\n", + " v = cutlass.Float32(0.0)\n", + " if col0 - 1 >= 0:\n", + " v = inp[row * W + (col0 - 1)]\n", + " tile[tx + 1, 0] = v\n", + " if ty == TS - 1:\n", + " v = cutlass.Float32(0.0)\n", + " if col0 + TS < W:\n", + " v = inp[row * W + (col0 + TS)]\n", + " tile[tx + 1, TS + 1] = v\n", + "\n", + " # Don't read the tile until every thread has finished filling it.\n", + " cute.arch.barrier()\n", + "\n", + " # Now the stencil is four on-chip lookups around the center cell.\n", + " acc = C0 * tile[tx + 1, ty + 1] + C1 * (\n", + " tile[tx, ty + 1] + tile[tx + 2, ty + 1] + tile[tx + 1, ty] + tile[tx + 1, ty + 2]\n", + " )\n", + " out[row * W + col] = acc" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## 3. Run both and verify\n", + "\n", + "A small factory bakes the chosen kernel into a launcher: one block per `TS x TS` output tile,\n", + "`block = (TS, TS, 1)`. We compare both kernels against a zero-padded NumPy-style reference; the\n", + "halo's zeros make the SMEM version agree with the naive one bit-for-bit (down to float epsilon)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "metadata": {}, + "outputs": [], + "source": [ + "def make_stencil(kernel):\n", + " @cute.jit\n", + " def run(inp: cutlass.Array, out: cutlass.Array, H: cutlass.Int32, W: cutlass.Int32):\n", + " grid = (H // TS, W // TS, 1)\n", + " kernel(inp, out, H, W).launch(grid=grid, block=(TS, TS, 1))\n", + " return run\n", + "\n", + "\n", + "def reference(x):\n", + " \"\"\"Host mirror: 5-point stencil with out-of-grid neighbors zeroed (zero-pad the grid).\"\"\"\n", + " xp = F.pad(x, (1, 1, 1, 1)) # one-cell zero border -> (H+2, W+2)\n", + " H, W = x.shape\n", + " north, south = xp[0:H, 1:W + 1], xp[2:H + 2, 1:W + 1]\n", + " west, east = xp[1:H + 1, 0:W], xp[1:H + 1, 2:W + 2]\n", + " return C0 * x + C1 * (north + south + west + east)\n", + "\n", + "\n", + "H, W = 64, 64\n", + "assert H % TS == 0 and W % TS == 0, \"H and W must be multiples of TS\"\n", + "x = torch.randn(H, W, dtype=torch.float32, device=\"cuda\")\n", + "ref = reference(x.cpu())\n", + "\n", + "for name, kern in ((\"naive\", stencil_naive_kernel), (\"smem\", stencil_smem_kernel)):\n", + " out = torch.zeros(H, W, dtype=torch.float32, device=\"cuda\")\n", + " make_stencil(kern)(cute.runtime.from_dlpack(x), cute.runtime.from_dlpack(out), H, W)\n", + " torch.testing.assert_close(out.cpu(), ref, atol=1e-4, rtol=1e-4)\n", + " print(f\"PASS {name}\")\n", + "\n", + "# Expected output:\n", + "# PASS naive\n", + "# PASS smem" + ] + }, + { + "cell_type": "markdown", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. **Change the physics.** Set `C0 = 1.0` and `C1 = 0.25` for a sharper diffusion step, or\n", + " `C0, C1 = 0.0, 0.25` for a pure neighbor-average (the center drops out). Update `reference`'s\n", + " coefficients to match and re-run — both kernels still agree.\n", + "2. **Iterate the heat equation.** Wrap the launch in a Python loop, ping-ponging two buffers\n", + " (`x -> out`, then `out -> x`), to watch a random field smooth out over N steps.\n", + "3. **Inspect the halo.** Temporarily set the halo loads to a sentinel (e.g. `-1.0`) instead of the\n", + " real neighbor, and print a border row of `out`: the corrupted values show exactly which cells\n", + " the halo feeds.\n", + "4. **Compare to `04_tiled_gemm`.** Both stage a tile into SMEM behind a `cute.arch.barrier()`; the\n", + " stencil's twist is the one-cell halo, since each output also needs its neighbors' data." + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "name": "python" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/3_kernels/03_blackwell_mma.ipynb b/examples/python/CuTeDSL/notebooks/3_kernels/03_blackwell_mma.ipynb new file mode 100644 index 0000000000..a4753154ff --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/3_kernels/03_blackwell_mma.ipynb @@ -0,0 +1,367 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Blackwell tcgen05 GEMM\n", + "\n", + "This is a complete warp-specialized `C = A @ B` for Blackwell's **5th-gen tensor cores**\n", + "(`tcgen05`). One CTA computes a single 128x128 output tile: a **TMA warp** streams the A and B\n", + "tiles from global memory into shared memory, an **MMA warp** drives `tcgen05_mma` into **Tensor\n", + "Memory (TMEM)**, and four **epilogue warps** read the result back out to C.\n", + "\n", + "The GEMM is **fp16 in, fp16 out, fp32 accumulate**: both A and B operands are `Float16`, the MMA\n", + "accumulates in `Float32` (in TMEM), and the epilogue casts the fp32 result back down to `Float16`\n", + "for C.\n", + "\n", + "The `tcgen05`/mbarrier intrinsics used here (the `prims.*` calls) are the **NVVM-level interface**\n", + "-- thin wrappers over the NVVM ops that drive the tensor cores, TMA engine, and mbarriers\n", + "directly. This notebook is also a tour of that interface.\n", + "\n", + "It is the smallest kernel that exercises the whole Blackwell GEMM pipeline -- TMA, mbarriers,\n", + "`tcgen05` MMA, and TMEM -- and the `04_blackwell_mma_contextvar` notebook scales exactly this shape\n", + "into a multi-CTA M x N x K GEMM." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** the three warp roles of a tcgen05 GEMM (TMA load, MMA, epilogue) and the two\n", + "mbarriers that chain them; how the result lands in **Tensor Memory (TMEM)** and how the epilogue\n", + "reads it back; and the **SMEM and instruction descriptors** that tell `tcgen05_mma` where its\n", + "operands sit and how to accumulate. **Prereq:** the `04_tiled_gemm` notebook (tiled GEMM, shared\n", + "memory, CTA barriers) and the `05_tma_load` notebook (TMA + mbarriers).\n", + "\n", + "**Runs on:** datacenter Blackwell -- `sm_100a` (B100/B200) and `sm_100f` (B300). `tcgen05` and TMEM\n", + "do not exist on Hopper or consumer `sm_120`. No GPU? It compiles without a GPU via\n", + "`CUTE_DSL_DRYRUN=1`." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import cutlass # Array, Float16/Float32, Int32/64, Constexpr, GridConstant, AddressSpace\n", + "import cutlass.cute as cute # @cute.kernel / @cute.jit, cute.arch.*, cute.runtime\n", + "from cutlass.experimental import primitives as prims # tcgen05 / mbarrier / TMA / barrier intrinsics + SMEM/Instr descriptors, TmemAddr\n", + "import cutlass.experimental.cuda as cuda # TMA descriptors: create_tensor_map_*, TensorMap(Swizzle)\n", + "\n", + "import torch\n", + "from typing import List" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. The warp-specialized skeleton\n", + "\n", + "Three warp roles split the work, chained by two mbarriers:\n", + "\n", + "```text\n", + " warp 0 TMA warp : copies A and B global -> SMEM, signals mbar_tma (operands ready)\n", + " warp 1 MMA warp : drives tcgen05_mma into Tensor Memory (TMEM), signals mbar_mma (result ready)\n", + " warps 2,3 (exit) : skipped, so the four epilogue warps form one clean warp-group\n", + " warps 4-7 epilogue : read the 128x128 result from TMEM and write it to C\n", + "```\n", + "\n", + "The result lands in **Tensor Memory (TMEM)** -- a Blackwell-only on-chip store that `tcgen05_mma`\n", + "writes and the epilogue reads back. The MMA reads its operands out of SMEM through two **SMEM\n", + "descriptors** (`Tcgen05SmemDesc`) whose swizzle and strides must match the TMA descriptor the host\n", + "builds. A third **instruction descriptor** (`Tcgen05InstrDesc`) sets the accumulation type: fp16\n", + "operands accumulate in fp32, which is why the epilogue reads fp32 out of TMEM and casts down to\n", + "fp16 for C." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "# =============================================================================\n", + "# Kernel: a warp-specialized tcgen05 GEMM for one 128x128 tile.\n", + "# TMA warp loads, MMA warp runs tcgen05_mma into TMEM, epilogue reads it.\n", + "# =============================================================================\n", + "@cute.kernel\n", + "def gemm_kernel(\n", + " tma_desc_a: cutlass.GridConstant[cuda.TensorMap],\n", + " tma_desc_b: cutlass.GridConstant[cuda.TensorMap],\n", + " matrix_c: cutlass.Array,\n", + " problem_size: cutlass.Constexpr[List[int]],\n", + ") -> None:\n", + " M, K, N = problem_size\n", + " thread_id, _, _ = cute.arch.thread_idx()\n", + " warp_id = cute.arch.warp_idx()\n", + " tmem_cols = N # the TMEM result tile is N columns wide\n", + "\n", + " # Step 1. Allocate the SMEM operand tiles, the two stage mbarriers, and a\n", + " # slot for the TMEM pointer.\n", + " smem_a = cutlass.Array(cutlass.Float16, (M, K), space=cutlass.AddressSpace.smem)\n", + " smem_b = cutlass.Array(cutlass.Float16, (N, K), space=cutlass.AddressSpace.smem)\n", + " mbar_tma = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem)\n", + " mbar_mma = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem)\n", + " tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem)\n", + "\n", + " # Step 2. One elected thread prefetches the TMA descriptors and arms both\n", + " # mbarriers (each expects a single arrival); the whole CTA syncs before use.\n", + " if prims.elect_sync():\n", + " prims.prefetch_tensormap(tma_desc_a.get_ptr())\n", + " prims.prefetch_tensormap(tma_desc_b.get_ptr())\n", + " prims.mbarrier_init(mbar_tma, 1)\n", + " prims.mbarrier_init(mbar_mma, 1)\n", + "\n", + " prims.fence_mbarrier_init()\n", + " cute.arch.barrier()\n", + "\n", + " # Step 3. Assign warp roles. Warps 2-3 are unused; exit them so the epilogue\n", + " # warp-group (4-7) lines up cleanly.\n", + " is_tma_warp = warp_id == 0\n", + " is_tc_warp = warp_id == 1\n", + " is_epi_warp = warp_id > 3\n", + " if warp_id == 2 or warp_id == 3:\n", + " prims.exit()\n", + "\n", + " # Step 4. The MMA warp owns the TMEM allocation; the result tile lives there\n", + " # until the epilogue drains it. Everyone reads the pointer after the alloc syncs.\n", + " if is_tc_warp:\n", + " prims.tcgen05_alloc(tmem_ptr_i32, tmem_cols)\n", + " cute.arch.barrier()\n", + " tmem_ptr = prims.make_tmem_ptr(tmem_ptr_i32.load(), cutlass.Int32)\n", + "\n", + " # Step 5. Run the warp-specialized body: TMA load, then MMA, then epilogue.\n", + " if is_tma_warp:\n", + " # Producer: trim the register file (light work), then have one thread\n", + " # post the expected byte count and kick off both TMA copies. The copies\n", + " # signal mbar_tma when the bytes land.\n", + " prims.setmaxregister(40, prims.SetMaxRegisterAction.DECREASE)\n", + " prims.bar_warp_sync(cute.arch.FULL_MASK)\n", + " if prims.elect_sync():\n", + " size_a = (K * M) * cutlass.Float16.width // 8\n", + " size_b = (K * N) * cutlass.Float16.width // 8\n", + " prims.mbarrier_arrive_expect_tx(mbar_tma, size_a + size_b)\n", + " prims.cp_async_bulk_tensor_shared_cta_global(\n", + " smem_a, tma_desc_a.get_ptr(), (0, 0), mbar_tma\n", + " )\n", + " prims.cp_async_bulk_tensor_shared_cta_global(\n", + " smem_b, tma_desc_b.get_ptr(), (0, 0), mbar_tma\n", + " )\n", + "\n", + " elif is_tc_warp:\n", + " # Consumer: wait for the operands, then run tcgen05_mma into TMEM.\n", + " prims.bar_warp_sync(cute.arch.FULL_MASK)\n", + " if prims.elect_sync():\n", + " while not prims.mbarrier_try_wait_parity(mbar_tma, 0):\n", + " pass\n", + " # SMEM descriptors describe where/how the operands sit in SMEM; the\n", + " # swizzle/strides must match what the host's TMA descriptor wrote\n", + " # (s128b -> stride_byte_offset 1024, leading_byte_offset 16).\n", + " desc_a = prims.Tcgen05SmemDesc.build(\n", + " smem_a,\n", + " leading_byte_offset=16,\n", + " stride_byte_offset=1024,\n", + " layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B,\n", + " )\n", + " desc_b = prims.Tcgen05SmemDesc.build(\n", + " smem_b,\n", + " leading_byte_offset=16,\n", + " stride_byte_offset=1024,\n", + " layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B,\n", + " )\n", + " # Instruction descriptor: fp16 operands accumulate in fp32.\n", + " idesc = prims.Tcgen05InstrDesc.build(\n", + " c_dtype=cutlass.Float32, n_dim=128, m_dim=128\n", + " )\n", + " # fp16 tensor-core K is 16, so K // 16 steps cover the K tile, each\n", + " # advancing the descriptors by 16 elements * 2 bytes. scale_d=False on\n", + " # the first step overwrites TMEM (no memset); True afterwards accumulates.\n", + " scale_d = False\n", + " for i in cutlass.range_constexpr(K // 16):\n", + " off = 16 * 2 * i\n", + " prims.tcgen05_mma(\n", + " prims.Tcgen05MMAKind.F16,\n", + " prims.CTAGroup.CTA_1,\n", + " tmem_ptr,\n", + " desc_a.advance_start_address(off),\n", + " desc_b.advance_start_address(off),\n", + " idesc,\n", + " scale_d,\n", + " )\n", + " scale_d = True\n", + " prims.tcgen05_commit(mbar_mma) # signal the result is ready\n", + "\n", + " elif is_epi_warp:\n", + " # Epilogue: wait for the result, then drain TMEM -> C. Each of the four\n", + " # warps owns a 32-row band of the 128-row tile. fp16 accumulated in fp32,\n", + " # so TMEM holds one fp32 per column; read two columns per thread and cast\n", + " # each down to fp16 for C.\n", + " while not prims.mbarrier_try_wait_parity(mbar_mma, 0):\n", + " pass\n", + " warpid_in_epi_wg = warp_id % 4\n", + " m_id = thread_id % 128\n", + " tmem_base = prims.TmemAddr(tmem_ptr_i32.load())\n", + " row_id = tmem_base.row_id + warpid_in_epi_wg * 32\n", + " for n in range(0, tmem_cols, 2):\n", + " tmem = prims.TmemAddr.from_row_col(\n", + " row_id, tmem_base.col_id + n\n", + " ).as_ptr(cutlass.Float32)\n", + " c_rmem = prims.tcgen05_ld(\"32x32b\", tmem, num=2)\n", + " for i in cutlass.range_constexpr(2):\n", + " matrix_c[m_id, n + i] = cutlass.Float16(c_rmem[i])\n", + "\n", + " cute.arch.barrier()\n", + "\n", + " # Step 6. Free the TMEM and release the allocation permit (MMA warp only).\n", + " if is_tc_warp:\n", + " prims.tcgen05_dealloc(tmem_ptr, tmem_cols)\n", + " prims.tcgen05_relinquish_alloc_permit()" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 2. Host: build the TMA descriptors and launch\n", + "\n", + "The host builds one tile-sized TMA descriptor per operand and launches a 256-thread (8-warp)\n", + "block. The descriptor **swizzle** must match the SMEM layout the kernel's `Tcgen05SmemDesc`\n", + "expects -- `s128b`, the 128-byte swizzle, paired with `stride_byte_offset=1024` in the kernel. A\n", + "flat `cutlass.Array` (what `from_dlpack` produces) feeds `create_tensor_map_tiled_from_view`\n", + "directly: it reads the source's shape, strides, and dtype to program the engine." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "# =============================================================================\n", + "# Host: build the A/B TMA descriptors and launch.\n", + "# =============================================================================\n", + "@cute.jit\n", + "def gemm(\n", + " matrix_a: cutlass.Array,\n", + " matrix_b: cutlass.Array,\n", + " matrix_c: cutlass.Array,\n", + " problem_size: cutlass.Constexpr[List[int]],\n", + ") -> None:\n", + " M, K, N = problem_size\n", + " # The TMA swizzle must match the kernel's SMEM layout (s128b). C is the\n", + " # epilogue output and stays a plain cutlass.Array (no descriptor).\n", + " tma_desc_a = cuda.create_tensor_map_tiled_from_view(\n", + " matrix_a, box_dims=(M, K), swizzle=cuda.TensorMapSwizzle.s128b\n", + " )\n", + " tma_desc_b = cuda.create_tensor_map_tiled_from_view(\n", + " matrix_b, box_dims=(N, K), swizzle=cuda.TensorMapSwizzle.s128b\n", + " )\n", + " # 256 threads = 8 warps, the split the kernel's warp roles assume.\n", + " gemm_kernel(tma_desc_a, tma_desc_b, matrix_c, problem_size).launch(\n", + " grid=(1, 1, 1), block=(256, 1, 1)\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "7", + "metadata": {}, + "source": [ + "## 3. Run it and check against PyTorch\n", + "\n", + "We multiply one 128x128 tile with K=64. A is `(M, K)` row-major; B is stored `(N, K)` (also\n", + "K-major), so the reference transposes it to form the `(M, K) @ (K, N)` product. The PyTorch\n", + "reference is computed on the **CPU**, so the check passes identically whether the kernel ran on real\n", + "hardware or compiled under dryrun." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "8", + "metadata": {}, + "outputs": [], + "source": [ + "# =============================================================================\n", + "# Main: run the GEMM and check against PyTorch.\n", + "# =============================================================================\n", + "# One 128x128 output tile with K=64.\n", + "M, N, K = 128, 128, 64\n", + "a = torch.randn(M, K, dtype=torch.float16, device=\"cuda\")\n", + "b = torch.randn(N, K, dtype=torch.float16, device=\"cuda\") # logical (N, K), K-major rows\n", + "c = torch.zeros(M, N, dtype=torch.float16, device=\"cuda\")\n", + "\n", + "gemm(\n", + " cute.runtime.from_dlpack(a),\n", + " cute.runtime.from_dlpack(b),\n", + " cute.runtime.from_dlpack(c),\n", + " (M, K, N),\n", + ")\n", + "\n", + "# CPU reference: B is stored (N, K), so transpose it to form (M, K) @ (K, N).\n", + "ref = a.cpu().float() @ b.cpu().float().T\n", + "torch.testing.assert_close(c.cpu().float(), ref, atol=1e-3, rtol=1e-3)\n", + "print(\"PASS\")\n", + "\n", + "# Expected output:\n", + "# PASS" + ] + }, + { + "cell_type": "markdown", + "id": "9", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. **The two mbarriers.** `mbar_tma` is signalled by the TMA copies (it counts *bytes*) and\n", + " `mbar_mma` by `tcgen05_commit` (it counts the MMA's completion). Trace which warp waits on which,\n", + " and why the epilogue can't read TMEM until `mbar_mma` fires.\n", + "2. **`scale_d`.** The first `tcgen05_mma` uses `scale_d=False` (overwrite TMEM), the rest `True`\n", + " (accumulate). What would change if every step used `True`, and why does \"first write overwrites\"\n", + " save a TMEM memset?\n", + "3. **The accumulator dtype.** The instruction descriptor sets `c_dtype=Float32`, so the epilogue\n", + " reads fp32 from TMEM and casts to fp16. Why must the read-back dtype match the accumulator the\n", + " MMA wrote?\n", + "4. **Descriptor swizzle.** The kernel's `Tcgen05SmemDesc` uses `SWIZZLE_128B` with\n", + " `stride_byte_offset=1024`, matching the host's `s128b`. What breaks if the host built the\n", + " descriptor with `swizzle=none` instead?" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/3_kernels/04_blackwell_mma_contextvar.ipynb b/examples/python/CuTeDSL/notebooks/3_kernels/04_blackwell_mma_contextvar.ipynb new file mode 100644 index 0000000000..c1702bcf90 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/3_kernels/04_blackwell_mma_contextvar.ipynb @@ -0,0 +1,467 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "0", + "metadata": {}, + "source": [ + "\n", + "\n", + "# Blackwell tcgen05 GEMM with contextvars\n", + "\n", + "The `03_blackwell_mma` notebook multiplied a single 128×128 tile. Here we scale that kernel into a\n", + "full `C = A @ B` over an arbitrary **M×N×K** problem, keeping the one-function-per-warp-role shape\n", + "that made the single-tile kernel readable.\n", + "\n", + "Two ideas carry the notebook:\n", + "\n", + "- **A 2-D grid of CTAs does the tiling.** Each CTA owns one 128×128 output tile, reads its\n", + " `(m, n)` tile coordinate from `cute.arch.block_idx()`, and walks the **K dimension** one 64-wide\n", + " tile at a time, accumulating in TMEM. One pair of *tile-sized* TMA descriptors serves the whole\n", + " grid: the per-CTA *coordinate*, not the descriptor, selects which slice of A and B to fetch.\n", + "- **Roles are still just Python.** Each warp picks its job with a plain `if warp_id == … elif …`\n", + " and runs its own branch. The light producer warp caps its register file with a\n", + " `with warp_registers(...)` block, so the heavy consumer and epilogue warps get a larger budget." + ] + }, + { + "cell_type": "markdown", + "id": "1", + "metadata": {}, + "source": [ + "**You'll learn:** how a grid of CTAs tiles a large GEMM (one 128×128 tile per CTA, addressed\n", + "by `cute.arch.block_idx()`); how one tile-sized TMA descriptor serves every CTA by steering\n", + "*coordinates* instead of rebuilding descriptors; how a K-loop accumulates across K-tiles in TMEM;\n", + "and the minimal producer↔consumer `full`/`empty` mbarrier handshake that lets a single SMEM buffer\n", + "be reused for every K-tile. The structure stays modular — one `@cute.jit` subroutine per role,\n", + "selected by a plain `if`/`elif` on the warp id — and a `with warp_registers(...)` block scopes the\n", + "producer warp's register budget.\n", + "\n", + "**Runs on:** **datacenter Blackwell `sm_100a` only** (B100/B200) — `tcgen05`/TMEM are not on\n", + "Hopper or consumer `sm_120`. No GPU? It compiles under dryrun (`CUTE_DSL_ARCH=sm_100a`)." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "2", + "metadata": {}, + "outputs": [], + "source": [ + "import contextlib\n", + "\n", + "import cutlass\n", + "import cutlass.cute as cute\n", + "from cutlass.experimental import primitives as prims # tcgen05 / mbarrier / TMA / barrier intrinsics + descriptors\n", + "import cutlass.experimental.cuda as cuda\n", + "import torch" + ] + }, + { + "cell_type": "markdown", + "id": "3", + "metadata": {}, + "source": [ + "## 1. Grid tiling, the K-loop, and the SMEM handshake\n", + "\n", + "The launch is a 2-D grid of `(N/128, M/128)` CTAs. Each CTA reads its tile coordinate from\n", + "`cute.arch.block_idx()` (`grid.x` → N tile, `grid.y` → M tile) and turns it into global row/column\n", + "offsets `(m_off, n_off)`. Inside the CTA the warps keep their roles:\n", + "\n", + "| warp(s) | role | per-CTA work |\n", + "|---|---|---|\n", + "| 0 | TMA producer | per K-tile, TMA-copies A[m_off, k] and B[k, n_off] global → SMEM |\n", + "| 1 | MMA consumer | per K-tile, runs `tcgen05_mma`; accumulates into the **same** TMEM tile |\n", + "| 2, 3 | exit | skipped, so the epilogue warps line up as a clean warp-group |\n", + "| 4-7 | epilogue | read the 128×128 result from TMEM and write C[m_off, n_off] |\n", + "\n", + "The new ingredient is the **K-loop** and its **single-buffer handshake**. We keep just one 128×64\n", + "SMEM tile per operand and reuse it for every K-tile, so the producer must not overwrite a tile the\n", + "consumer is still reading. Two mbarriers chain them into a ping-pong:\n", + "\n", + "| mbarrier | direction | meaning |\n", + "|---|---|---|\n", + "| `mbar_full` | producer → consumer | \"this K-tile is loaded, go.\" |\n", + "| `mbar_empty` | consumer → producer | \"I'm done with the buffer, refill it.\" Pre-signalled once before the loop so the first load doesn't deadlock on a consumer that hasn't run. |\n", + "\n", + "With exactly one buffer, the handshake alternates between two mbarrier *phases*, so each side waits\n", + "on parity `kt % 2`. The consumer accumulates with `scale_d = kt != 0`: the first MMA overwrites\n", + "TMEM (no memset needed), every later one adds to it. When the K reduction finishes, `mbar_mma`\n", + "signals the epilogue. The TMA *coordinates* are `(k_off, m_off)` for A and `(k_off, n_off)` for B —\n", + "both operands are stored K-major, so K leads.\n", + "\n", + "> **Metaprogramming aside.** `warp_registers` (next cell) is an ordinary Python **context manager** -- it scopes a *policy* with a `with` block, no GPU state involved. The same Python tool scales to launch configuration: a module-level `contextvars.ContextVar` set by a `with launch_config(block=...)` block lets a caller pick the launch shape without threading it through every call. Both keep the kernel body clean by moving configuration into plain Python." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "4", + "metadata": {}, + "outputs": [], + "source": [ + "# Output tile is 128×128; K is consumed 64 elements at a time (64 fp16 = 128 B,\n", + "# which is exactly the s128b swizzle box).\n", + "TILE_M = 128\n", + "TILE_N = 128\n", + "TILE_K = 64\n", + "\n", + "\n", + "@contextlib.contextmanager\n", + "def warp_registers(regs, action=prims.SetMaxRegisterAction.DECREASE):\n", + " \"\"\"Scope a per-warp register budget with a `with` block.\n", + "\n", + " This is pure register *policy* (`setmaxregister`) — it changes how many\n", + " registers the warp may use, not which threads run. The surrounding\n", + " `if warp_id == ...` already selects the warp; this block only sets the budget.\n", + " \"\"\"\n", + " prims.setmaxregister(regs, action)\n", + " yield\n", + "\n", + "\n", + "# =============================================================================\n", + "# Warp workers: one @cute.jit subroutine per role.\n", + "# =============================================================================\n", + "@cute.jit\n", + "def tma_producer(\n", + " smem_a,\n", + " smem_b,\n", + " tma_desc_a,\n", + " tma_desc_b,\n", + " mbar_full,\n", + " mbar_empty,\n", + " m_off,\n", + " n_off,\n", + " num_k_tiles,\n", + "):\n", + " # One elected thread drives the whole warp's TMA traffic.\n", + " prims.bar_warp_sync(cute.arch.FULL_MASK)\n", + " if prims.elect_sync():\n", + " size_a = (TILE_K * TILE_M) * cutlass.Float16.width // 8\n", + " size_b = (TILE_K * TILE_N) * cutlass.Float16.width // 8\n", + " for kt in range(num_k_tiles):\n", + " # Step 1. Wait until the consumer has freed the SMEM buffer.\n", + " # `empty` is pre-signalled (phase 1) and the consumer flips it once\n", + " # per K-tile, so the phase to wait on before iteration kt is kt % 2 —\n", + " # which lets the first iteration pass straight through.\n", + " while not prims.mbarrier_try_wait_parity(mbar_empty, kt % 2):\n", + " pass\n", + " # Step 2. Fire the two bulk TMA copies for this CTA's (m, n) tile.\n", + " # A is (M, K), B is (K, N) — both K-major — so the TMA coordinate\n", + " # leads with the K offset.\n", + " k_off = kt * TILE_K\n", + " prims.mbarrier_arrive_expect_tx(mbar_full, size_a + size_b)\n", + " prims.cp_async_bulk_tensor_shared_cta_global(\n", + " smem_a, tma_desc_a.get_ptr(), (k_off, m_off), mbar_full\n", + " )\n", + " prims.cp_async_bulk_tensor_shared_cta_global(\n", + " smem_b, tma_desc_b.get_ptr(), (k_off, n_off), mbar_full\n", + " )\n", + "\n", + "\n", + "@cute.jit\n", + "def mma_consumer(\n", + " smem_a, smem_b, tmem_ptr, mbar_full, mbar_empty, mbar_mma, num_k_tiles\n", + "):\n", + " # One elected thread issues the MMAs for the whole warp.\n", + " prims.bar_warp_sync(cute.arch.FULL_MASK)\n", + " if prims.elect_sync():\n", + " idesc = prims.Tcgen05InstrDesc.build(\n", + " c_dtype=cutlass.Float32, n_dim=TILE_N, m_dim=TILE_M\n", + " )\n", + " for kt in range(num_k_tiles):\n", + " # Step 1. Wait for the producer's data for this K-tile.\n", + " while not prims.mbarrier_try_wait_parity(mbar_full, kt % 2):\n", + " pass\n", + " # Step 2. Run tcgen05 MMA into TMEM, accumulating across K.\n", + " desc_a = prims.Tcgen05SmemDesc.build(\n", + " smem_a,\n", + " leading_byte_offset=16,\n", + " stride_byte_offset=1024,\n", + " layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B,\n", + " )\n", + " desc_b = prims.Tcgen05SmemDesc.build(\n", + " smem_b,\n", + " leading_byte_offset=16,\n", + " stride_byte_offset=1024,\n", + " layout=prims.Tcgen05SmemSwizzle.SWIZZLE_128B,\n", + " )\n", + " # Only the first MMA of the whole K-loop overwrites TMEM; the rest\n", + " # accumulate on top. `kt` is a staged int, so compare against 0.\n", + " scale_d = kt != 0\n", + " # fp16 tensor-core K is 16, so 4 steps cover one 64-wide K-tile.\n", + " for i in cutlass.range_constexpr(4):\n", + " off: cutlass.Constexpr[int] = 16 * 2 * i\n", + " prims.tcgen05_mma(\n", + " prims.Tcgen05MMAKind.F16,\n", + " prims.CTAGroup.CTA_1,\n", + " tmem_ptr,\n", + " desc_a.advance_start_address(off),\n", + " desc_b.advance_start_address(off),\n", + " idesc,\n", + " scale_d,\n", + " )\n", + " scale_d = True\n", + " # Step 3. Release the SMEM buffer back to the producer.\n", + " prims.tcgen05_commit(mbar_empty)\n", + " # Step 4. Whole accumulation done: wake the epilogue.\n", + " prims.tcgen05_commit(mbar_mma)\n", + "\n", + "\n", + "@cute.jit\n", + "def epilogue(matrix_c, tmem_ptr_i32, mbar_mma, thread_id, warp_id, m_off, n_off, tmem_num_col):\n", + " # Step 1. Wait for the MMA to finish.\n", + " while not prims.mbarrier_try_wait_parity(mbar_mma, 0):\n", + " pass\n", + " # Step 2. Copy the TMEM accumulator out to this CTA's 128×128 tile of C\n", + " # at global offset (m_off, n_off).\n", + " warpid_in_epi_wg = warp_id % 4\n", + " tmem_base = prims.TmemAddr(tmem_ptr_i32.load())\n", + " row_id = tmem_base.row_id + warpid_in_epi_wg * 32\n", + " for n in range(0, tmem_num_col, 2):\n", + " tmem = prims.TmemAddr.from_row_col(row_id, tmem_base.col_id + n).as_ptr(\n", + " cutlass.Float32\n", + " )\n", + " c_rmem = prims.tcgen05_ld(\"32x32b\", tmem, num=2)\n", + " m_id = thread_id % 128\n", + " for i in cutlass.range_constexpr(2):\n", + " matrix_c[m_off + m_id, n_off + n + i] = cutlass.Float16(c_rmem[i])\n", + "\n", + "\n", + "# =============================================================================\n", + "# Kernel: set up shared state, then dispatch each warp to its worker.\n", + "# =============================================================================\n", + "@cute.kernel\n", + "def gemm_kernel(\n", + " tma_desc_a: cutlass.GridConstant[cuda.TensorMap],\n", + " tma_desc_b: cutlass.GridConstant[cuda.TensorMap],\n", + " matrix_c: cutlass.Array,\n", + " problem_size: cutlass.Constexpr,\n", + ") -> None:\n", + " # Step 1. Read this CTA's output tile from its block index: grid.x indexes\n", + " # N tiles, grid.y indexes M tiles.\n", + " M, K, N = problem_size\n", + " num_k_tiles: cutlass.Constexpr = K // TILE_K\n", + " bid_n, bid_m, _ = cute.arch.block_idx()\n", + " m_off = bid_m * TILE_M\n", + " n_off = bid_n * TILE_N\n", + " thread_id, _, _ = cute.arch.thread_idx()\n", + " warp_id = cute.arch.warp_idx()\n", + " tmem_num_col = 128\n", + "\n", + " # Step 2. Declare shared state: one 128×64 buffer per operand (reused across\n", + " # K-tiles via the full/empty handshake), the three mbarriers, and a slot for\n", + " # the TMEM pointer.\n", + " smem_a = cutlass.Array(\n", + " cutlass.Float16, (TILE_M, TILE_K), space=cutlass.AddressSpace.smem\n", + " )\n", + " smem_b = cutlass.Array(\n", + " cutlass.Float16, (TILE_N, TILE_K), space=cutlass.AddressSpace.smem\n", + " )\n", + " mbar_full = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem)\n", + " mbar_empty = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem)\n", + " mbar_mma = cutlass.Array(cutlass.Int64, 1, space=cutlass.AddressSpace.smem)\n", + " tmem_ptr_i32 = cutlass.Array(cutlass.Int32, 1, space=cutlass.AddressSpace.smem)\n", + "\n", + " # Step 3. One elected thread prefetches the descriptors and inits the\n", + " # mbarriers. `empty` is pre-armed so the producer's first K-tile can proceed.\n", + " if prims.elect_sync():\n", + " prims.prefetch_tensormap(tma_desc_a.get_ptr())\n", + " prims.prefetch_tensormap(tma_desc_b.get_ptr())\n", + " prims.mbarrier_init(mbar_full, 1)\n", + " prims.mbarrier_init(mbar_empty, 1)\n", + " prims.mbarrier_init(mbar_mma, 1)\n", + " prims.mbarrier_arrive(mbar_empty) # pre-signal: buffer starts free\n", + " prims.fence_mbarrier_init()\n", + " cute.arch.barrier()\n", + "\n", + " # Step 4. Retire the unused warps. Warps 2-3 exit here, leaving 4-7 as a\n", + " # clean warp-group for the epilogue. The CTA-wide barriers stay at kernel\n", + " # scope, *between* the role `if` blocks, so every active warp reaches them.\n", + " TMA_WARP, MMA_WARP = 0, 1\n", + "\n", + " if warp_id == 2 or warp_id == 3:\n", + " prims.exit()\n", + "\n", + " # Step 5. The consumer warp owns the TMEM accumulator; allocate it before\n", + " # anyone uses it.\n", + " if warp_id == MMA_WARP:\n", + " prims.tcgen05_alloc(tmem_ptr_i32, tmem_num_col)\n", + " cute.arch.barrier()\n", + " tmem_ptr = prims.make_tmem_ptr(tmem_ptr_i32.load(), cutlass.Int32)\n", + "\n", + " # Step 6. Dispatch each warp to its worker. Producer and consumer run\n", + " # concurrently, synchronized by the SMEM full/empty mbarriers.\n", + " if warp_id == TMA_WARP:\n", + " # The producer is light, so cap its register budget for this scope. The\n", + " # `if` above already selected the warp; the context manager only sets\n", + " # policy, so it adds no gating.\n", + " with warp_registers(40):\n", + " tma_producer(\n", + " smem_a,\n", + " smem_b,\n", + " tma_desc_a,\n", + " tma_desc_b,\n", + " mbar_full,\n", + " mbar_empty,\n", + " m_off,\n", + " n_off,\n", + " num_k_tiles,\n", + " )\n", + " elif warp_id == MMA_WARP:\n", + " mma_consumer(\n", + " smem_a, smem_b, tmem_ptr, mbar_full, mbar_empty, mbar_mma, num_k_tiles\n", + " )\n", + " elif warp_id == 4 or warp_id == 5 or warp_id == 6 or warp_id == 7:\n", + " epilogue(\n", + " matrix_c, tmem_ptr_i32, mbar_mma, thread_id, warp_id, m_off, n_off, tmem_num_col\n", + " )\n", + "\n", + " # Step 7. Sync, then the consumer releases the TMEM it allocated.\n", + " cute.arch.barrier()\n", + " if warp_id == MMA_WARP:\n", + " prims.tcgen05_dealloc(tmem_ptr, tmem_num_col)\n", + " prims.tcgen05_relinquish_alloc_permit()" + ] + }, + { + "cell_type": "markdown", + "id": "5", + "metadata": {}, + "source": [ + "## 2. Host: one descriptor pair, a 2-D grid of tiles\n", + "\n", + "The trick: the TMA descriptors span the **whole** A and B but carry a **tile-sized** 128×64 box.\n", + "That single pair covers every CTA — the per-CTA coordinate slides the box to the right tile, so we\n", + "never rebuild a descriptor per CTA. A is row-major `(M, K)`; B is K-major `(K, N)`, so its box and\n", + "TMA coordinates lead with K to match A. The launch is a 2-D grid `(N/128, M/128)` of 256-thread\n", + "(8-warp) blocks — one block per 128×128 output tile." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "6", + "metadata": {}, + "outputs": [], + "source": [ + "# =============================================================================\n", + "# Host: build the A/B TMA descriptors and launch the multi-CTA GEMM kernel.\n", + "# =============================================================================\n", + "@cute.jit\n", + "def gemm(\n", + " matrix_a: cutlass.Array,\n", + " matrix_b: cutlass.Array,\n", + " matrix_c: cutlass.Array,\n", + " problem_size: cutlass.Constexpr,\n", + ") -> None:\n", + " M, K, N = problem_size\n", + " # Tile-sized descriptors over the WHOLE A and B: the box is one 128×64 tile,\n", + " # and each CTA steers it with coordinates, so this single pair serves the grid.\n", + " tma_desc_a = cuda.create_tensor_map_tiled_from_view(\n", + " matrix_a, box_dims=(TILE_M, TILE_K), swizzle=cuda.TensorMapSwizzle.s128b\n", + " )\n", + " # B is stored K-major (shape (K, N), K contiguous), so its mode order is\n", + " # (K, N): box and TMA coords lead with K, matching A's convention.\n", + " tma_desc_b = cuda.create_tensor_map_tiled_from_view(\n", + " matrix_b, box_dims=(TILE_K, TILE_N), swizzle=cuda.TensorMapSwizzle.s128b\n", + " )\n", + " # A 2-D grid of 128×128 output tiles: grid.x over N, grid.y over M.\n", + " grid_n = N // TILE_N\n", + " grid_m = M // TILE_M\n", + " gemm_kernel(tma_desc_a, tma_desc_b, matrix_c, problem_size).launch(\n", + " grid=(grid_n, grid_m, 1), block=(256, 1, 1)\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "7", + "metadata": {}, + "source": [ + "## 3. Run it and check against PyTorch\n", + "\n", + "The driver runs `M = N = 256`, `K = 128` — a **2×2 grid** of 128×128 tiles, each CTA doing a\n", + "2-step K-loop. `b` is built transposed (`.T` flips only the *logical* shape, leaving the data\n", + "K-major), so it matches the descriptor's `(K, N)` layout. The PyTorch reference is computed on the\n", + "CPU, so the correctness check passes identically on real Blackwell hardware." + ] + }, + { + "cell_type": "code", + "execution_count": null, + "id": "8", + "metadata": {}, + "outputs": [], + "source": [ + "# =============================================================================\n", + "# Main: build A/B/C, run the GEMM, and verify against PyTorch.\n", + "# =============================================================================\n", + "# Step 1. Build A/B/C. A is (M, K) row-major; B is (N, K) but stored K-major via\n", + "# .T. At 256×256×128 this is a 2×2 grid of 128×128 tiles, each CTA running a\n", + "# 2-step K-loop.\n", + "M, N, K = 256, 256, 128\n", + "a = torch.randn(M, K, dtype=torch.float16, device=\"cuda\")\n", + "b = torch.randn(N, K, dtype=torch.float16, device=\"cuda\").T # K-major\n", + "c = torch.zeros(M, N, dtype=torch.float16, device=\"cuda\")\n", + "\n", + "# Step 2. Run the GEMM.\n", + "gemm(\n", + " cute.runtime.from_dlpack(a),\n", + " cute.runtime.from_dlpack(b),\n", + " cute.runtime.from_dlpack(c),\n", + " (M, K, N),\n", + ")\n", + "\n", + "# Step 3. Verify. Reference on the CPU so the correctness check runs identically\n", + "# on real Blackwell hardware.\n", + "ref = a.cpu().float() @ b.cpu().float()\n", + "torch.testing.assert_close(c.cpu().float(), ref, atol=1e-2, rtol=1e-2)\n", + "print(\"PASS\")\n", + "\n", + "# Expected output:\n", + "# PASS" + ] + }, + { + "cell_type": "markdown", + "id": "9", + "metadata": {}, + "source": [ + "## Try it yourself\n", + "\n", + "1. The grid is `(N/128, M/128)` and `bid_n, bid_m, _ = cute.arch.block_idx()`. Why is `grid.x`\n", + " mapped to N and `grid.y` to M? What breaks if M or N is not a multiple of 128?\n", + "2. `mbar_empty` is pre-signalled once before the loop. Trace the `full`/`empty` parities across\n", + " `kt = 0, 1, 2, …` — why must the producer wait on `kt % 2` (not `(kt+1) % 2`)?\n", + "3. `scale_d = kt != 0` makes only the first K-tile overwrite TMEM. What would memset-ing the\n", + " accumulator instead cost, and why is \"first write overwrites\" cheaper for a long K-loop?\n", + "4. Every CTA loads its own A row-block and B col-block independently. Which loads do neighbouring\n", + " CTAs *share*, and how would CTA-multicast TMA (`...shared_cluster_global`) exploit that?" + ] + } + ], + "metadata": { + "kernelspec": { + "display_name": "Python 3 (ipykernel)", + "language": "python", + "name": "python3" + }, + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.3" + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/examples/python/CuTeDSL/notebooks/README.md b/examples/python/CuTeDSL/notebooks/README.md new file mode 100644 index 0000000000..50eb6be303 --- /dev/null +++ b/examples/python/CuTeDSL/notebooks/README.md @@ -0,0 +1,46 @@ +# CuTeDSL tutorial-notebook course + +A hands-on Jupyter-notebook course for the **CuTeDSL** Primitives programming model. +Each chapter is a self-contained notebook that runs end-to-end on a CUDA GPU. + +## Chapters + +### 1. DSL features — [`1_dsl_features/`](1_dsl_features) +| notebook | topic | +|---|---| +| `01_hello_world.ipynb` | `@cute.jit` / `@cute.kernel`; `print()` (host, trace-time) vs `cute.printf()` (device, per-thread) | +| `02_control_flow.ipynb` | meta loops vs staged loops — `range_constexpr` / `range` / `range(unroll=)`, `const_expr` branches | +| `03_diagnostics.ipynb` | compiler diagnostics — the `warnings{...}` / `remarks{...}` options | +| `04_zero_cost_abstraction.ipynb` | classes & polymorphism are trace-time Python — they compile away (PTX proof) | + +### 2. Primitives — [`2_primitives/`](2_primitives) +| notebook | topic | +|---|---| +| `01_array_concepts.ipynb` | `cutlass.Array` and the GPU memory spaces | +| `02_vector_concepts.ipynb` | `cutlass.Vector` — registers vs memory | +| `03_cute_interop.ipynb` | `cute.Tensor` ↔ `cutlass.Array` interop | +| `04_tiled_gemm.ipynb` | tiled GEMM — `Array`, shared memory, passing a `Callable` as kernel metaprogramming | +| `05_tma_load.ipynb` | TMA load + `mbarrier` | +| `06_prefix_sum.ipynb` | prefix sum (warp-shuffle scan) | + +### 3. Kernels — [`3_kernels/`](3_kernels) +| notebook | topic | +|---|---| +| `01_softmax.ipynb` | softmax (naive → online/flash) | +| `02_stencil_2d.ipynb` | 2D 5-point stencil (a heat-diffusion step) — naive vs shared-memory halo tiling | +| `03_blackwell_mma.ipynb` | minimal Blackwell `tcgen05` GEMM (sm_100a) | +| `04_blackwell_mma_contextvar.ipynb` | the same `tcgen05` GEMM driven by Python `contextvars` (sm_100a) | + +## Related notebooks + +The long-standing CuTe-layout and CUDA-runtime series — CuTe layout algebra, async +pipelines, composed layouts, CUDA graphs, autotuning, and the tour to a +state-of-the-art GEMM — live alongside this course under +[`../cute/notebooks/`](../cute/notebooks). + +## Running the notebooks + +Open any notebook in Jupyter on a CUDA host. A few chapters target Blackwell +(`sm_100a`) tensor cores (`3_kernels/03`, `04`) and need a B100/B200-class GPU; the +rest run on any recent CUDA GPU. The course chapters are also executed end-to-end in +CI via `nbconvert`. diff --git a/include/cute/arch/mma_sm107_umma.hpp b/include/cute/arch/mma_sm107_umma.hpp index eaa97ba67f..20adf48b08 100644 --- a/include/cute/arch/mma_sm107_umma.hpp +++ b/include/cute/arch/mma_sm107_umma.hpp @@ -333,9 +333,9 @@ struct SM107_MMA_MXF4NVF4_SS static_assert((VS == 16) || (VS == 32), "SM107_MMA_MXF4NVF4_SS Vector size can only be 16 or 32."); static_assert(is_same_v || - is_same_v || + is_same_v || is_same_v, - "SF data type can only be one of {ue8m0, e4m3, or ue5m3}."); + "SF data type can only be one of {ue8m0, ue4m3, or ue5m3}."); using DRegisters = void; using ARegisters = uint64_t[1]; @@ -449,9 +449,9 @@ struct SM107_MMA_MXF4NVF4_2x1SM_SS static_assert((VS == 16) || (VS == 32), "SM107_MMA_MXF4NVF4_2x1SM_SS Vector size can only be 16 or 32."); static_assert(is_same_v || - is_same_v || + is_same_v || is_same_v, - "SF data type can only be one of {ue8m0, e4m3, or ue5m3}."); + "SF data type can only be one of {ue8m0, ue4m3, or ue5m3}."); using DRegisters = void; using ARegisters = uint64_t[1]; diff --git a/include/cutlass/arch/barrier.h b/include/cutlass/arch/barrier.h index 519005e507..35c7763b34 100644 --- a/include/cutlass/arch/barrier.h +++ b/include/cutlass/arch/barrier.h @@ -386,6 +386,12 @@ struct ClusterBarrier { ClusterBarrier::arrive(&this->barrier_, cta_id, pred); } + // Remote SMEM arrive with relaxed semantics and cluster scope + CUTLASS_DEVICE + void arrive_relaxed_cluster(uint32_t cta_id, uint32_t pred = true) const { + ClusterBarrier::arrive_relaxed_cluster(&this->barrier_, cta_id, pred); + } + // // Static Versions // @@ -506,6 +512,29 @@ struct ClusterBarrier { #endif } + // Same as the above, with relaxed sem and cluster scope. + CUTLASS_HOST_DEVICE + static void arrive_relaxed_cluster(ValueType const* smem_ptr, uint32_t cta_id, uint32_t pred) { +#if CUDA_BARRIER_ENABLED + uint32_t smem_addr = cute::cast_smem_ptr_to_uint(smem_ptr); + if (pred) { + asm volatile( + "{\n\t" + ".reg .b32 remAddr32;\n\t" + "mapa.shared::cluster.u32 remAddr32, %0, %1;\n\t" + "mbarrier.arrive.relaxed.cluster.shared::cluster.b64 _, [remAddr32];\n\t" + "}" + : + : "r"(smem_addr), "r"(cta_id) + : "memory"); + } + + cutlass::arch::synclog_emit_cluster_barrier_arrive_cluster(__LINE__, smem_addr, cta_id, pred); +#else + CUTLASS_NOT_IMPLEMENTED(); +#endif + } + // Barrier arrive on local smem CUTLASS_HOST_DEVICE static void arrive(ValueType const* smem_ptr) { @@ -562,6 +591,12 @@ struct ClusterTransactionBarrier : public ClusterBarrier { ClusterTransactionBarrier::arrive_and_expect_tx(&this->barrier_, transaction_bytes , cta_id, pred); } + // Same as above, with relaxed semantics and cluster scope + CUTLASS_DEVICE + void arrive_and_expect_tx_relaxed_cluster(uint32_t transaction_bytes, uint32_t cta_id, uint32_t pred = 1u) const { + ClusterTransactionBarrier::arrive_and_expect_tx_relaxed_cluster(&this->barrier_, transaction_bytes, cta_id, pred); + } + // Performs an expected transaction bytes increment without doing an arrive operation CUTLASS_DEVICE void expect_transaction(uint32_t transaction_bytes) const { @@ -625,6 +660,28 @@ struct ClusterTransactionBarrier : public ClusterBarrier { #endif } + // Same as the above, with relaxed sem and cluster scope. + CUTLASS_HOST_DEVICE + static void arrive_and_expect_tx_relaxed_cluster( + ValueType const* smem_ptr, uint32_t transaction_bytes, uint32_t cta_id, uint32_t pred) { +#if CUDA_BARRIER_ENABLED + uint32_t smem_addr = cute::cast_smem_ptr_to_uint(smem_ptr); + asm volatile( + "{\n\t" + ".reg .pred p;\n\t" + ".reg .b32 remAddr32;\n\t" + "setp.eq.u32 p, %2, 1;\n\t" + "@p mapa.shared::cluster.u32 remAddr32, %0, %1;\n\t" + "@p mbarrier.arrive.expect_tx.relaxed.cluster.shared::cluster.b64 _, [remAddr32], %3;\n\t" + "}" + : + : "r"(smem_addr), "r"(cta_id), "r"(pred), "r"(transaction_bytes) + : "memory"); +#else + CUTLASS_NOT_IMPLEMENTED(); +#endif + } + // Performs an expected transaction bytes increment without doing an arrive operation CUTLASS_HOST_DEVICE static void expect_transaction(ValueType const* smem_ptr, uint32_t transaction_bytes) { @@ -948,7 +1005,6 @@ CUTE_DEVICE static void fence_view_async_tmem_store() { #endif } - //////////////////////////////////////////////////////////////////////////////////////////////////// } // end namespace arch } // end namespace cutlass diff --git a/include/cutlass/array.h b/include/cutlass/array.h index 0ee2591f98..a7f98152fe 100644 --- a/include/cutlass/array.h +++ b/include/cutlass/array.h @@ -546,6 +546,15 @@ Array make_Array(Element x, Element y, Element z, Element w) { // functional.h numeric specializations ///////////////////////////////////////////////////////////////////////////////////////////////// +template +struct no_op< Array > { + + CUTLASS_HOST_DEVICE + Array operator()(Array const &lhs) const { + return lhs; + } +}; + template struct absolute_value_op< Array > { diff --git a/include/cutlass/detail/layout.hpp b/include/cutlass/detail/layout.hpp index f3eada3bb2..032ae434d6 100644 --- a/include/cutlass/detail/layout.hpp +++ b/include/cutlass/detail/layout.hpp @@ -37,7 +37,6 @@ #include "cute/util/type_traits.hpp" #include "cute/arch/copy_sm90_tma.hpp" #include "cute/arch/copy_sm100_tma.hpp" - #include "cutlass/layout/matrix.h" #include "cutlass/layout/tensor.h" #include "cutlass/numeric_types.h" @@ -302,17 +301,17 @@ constexpr bool is_tma_copy_engine() { return false; } else { - if constexpr ( cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v - || cute::is_base_of_v + if constexpr ( cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v + || cute::is_base_of_v ) { return true; } diff --git a/include/cutlass/detail/sm107_blockscaled_layout.hpp b/include/cutlass/detail/sm107_blockscaled_layout.hpp index 865380addc..e25486fe6c 100644 --- a/include/cutlass/detail/sm107_blockscaled_layout.hpp +++ b/include/cutlass/detail/sm107_blockscaled_layout.hpp @@ -136,14 +136,22 @@ struct Sm107BlockScaledConfig { CUTE_HOST_DEVICE static constexpr auto deduce_smem_layoutSFB(TiledMma tiled_mma, TileShape_MNK tileshape_mnk) { - // CTA-level MMA tile shape (TILE_N, TILE_K) + // CTA-level MMA tile shape (Round_Up(TILE_N, 128), TILE_K) + // This round_up workaround is used for TileShape_N == 64 or 192, which don't + // natively divide the SF atom's 128-wide MN mode; + // Copies (from GMEM to SMEM, and SMEM to TMEM) transfer SFB as if N were + // rounded up to 128/256. constexpr auto tile_shape_nk = - make_shape(size<1>(TileShape_MNK{}), size<2>(TileShape_MNK{})); + make_shape(Int(TileShape_MNK{}), Blk_MN{}) * Blk_MN{}>{}, + size<2>(TileShape_MNK{})); - // CTA-level MMA instruction shape (MMA_INST_N, MMA_INST_K) - constexpr auto mma_shape_nk = - make_shape(size<1>(typename TiledMma::AtomShape_MNK{}), - size<2>(typename TiledMma::AtomShape_MNK{})); + // Number of MMA instructions needed to cover the CTA tile + constexpr auto mma_tile_inst_nk = + make_shape(size<1>(TileShape_MNK{}) / size<1>(typename TiledMma::AtomShape_MNK{}), + size<2>(TileShape_MNK{}) / size<2>(typename TiledMma::AtomShape_MNK{})); + + // Effective per-instruction atom size within the rounded-up SF tile (MMA_INST_N, MMA_INST_K) + constexpr auto mma_shape_nk = shape_div(tile_shape_nk, mma_tile_inst_nk); // Tiling the CTA-level tile shape with SF atoms, first accross the K-mode, and then N-mode auto smem_layout_tiled = tile_to_shape(SfAtom{}, tile_shape_nk, Step<_2, _1>{}); diff --git a/include/cutlass/epilogue/collective/builders/sm107_builder.inl b/include/cutlass/epilogue/collective/builders/sm107_builder.inl index bc47f473a5..6d5b5eb0d8 100644 --- a/include/cutlass/epilogue/collective/builders/sm107_builder.inl +++ b/include/cutlass/epilogue/collective/builders/sm107_builder.inl @@ -81,7 +81,7 @@ struct CollectiveBuilder< GmemLayoutTagD, AlignmentD, EpilogueScheduleType, - FusionOp + FusionOp > { using CollectiveOp = typename CollectiveBuilder< @@ -103,4 +103,6 @@ struct CollectiveBuilder< >::CollectiveOp; }; +///////////////////////////////////////////////////////////////////////////////////////////////// + } // namespace cutlass::epilogue::collective diff --git a/include/cutlass/epilogue/fusion/operations.hpp b/include/cutlass/epilogue/fusion/operations.hpp index 1ee68514ab..cc08532709 100644 --- a/include/cutlass/epilogue/fusion/operations.hpp +++ b/include/cutlass/epilogue/fusion/operations.hpp @@ -90,6 +90,18 @@ struct FusionOperation { using GmemLayoutTagScalefactor = void; }; +// D = acc +template< + class ElementOutput_, + class ElementCompute_, + FloatRoundStyle RoundStyle_ = FloatRoundStyle::round_to_nearest +> +struct PlainAcc : FusionOperation { + using ElementOutput = ElementOutput_; + using ElementCompute = ElementCompute_; + static constexpr auto RoundStyle = RoundStyle_; +}; + // D = alpha * acc template< class ElementOutput_, diff --git a/include/cutlass/epilogue/fusion/sm90_callbacks_tma_warpspecialized.hpp b/include/cutlass/epilogue/fusion/sm90_callbacks_tma_warpspecialized.hpp index a0cfe55b33..d02f9bc671 100644 --- a/include/cutlass/epilogue/fusion/sm90_callbacks_tma_warpspecialized.hpp +++ b/include/cutlass/epilogue/fusion/sm90_callbacks_tma_warpspecialized.hpp @@ -57,6 +57,44 @@ namespace cutlass::epilogue::fusion { template using Sm90EVT = Sm90TreeVisitor; +// D = acc +template < + int StagesC, + int StagesD, + int FragmentSize, + bool ReuseSmemC, + bool DelayTmaStore, + class ElementOutput, + class ElementCompute, + FloatRoundStyle RoundStyle, + class CtaTileShapeMNK, + class EpilogueTile +> +struct FusionCallbacks< + epilogue::Sm90TmaWarpSpecialized, + fusion::PlainAcc, + CtaTileShapeMNK, + EpilogueTile +> : Sm90EVT, + Sm90AccFetch + > { + using Impl = + Sm90EVT, + Sm90AccFetch + >; + using Operation = fusion::PlainAcc; + + struct Arguments { + operator typename Impl::Arguments() const { + return + {}; + } + }; + + // Ctor inheritance + using Impl::Impl; +}; + // D = alpha * acc template < int StagesC, diff --git a/include/cutlass/functional.h b/include/cutlass/functional.h index 2ac6e6539c..eb9d20b984 100644 --- a/include/cutlass/functional.h +++ b/include/cutlass/functional.h @@ -105,6 +105,14 @@ namespace detail { ///////////////////////////////////////////////////////////////////////////////////////////////// +template +struct no_op { + CUTLASS_HOST_DEVICE + T operator()(T lhs) const { + return lhs; + } +}; + template struct absolute_value_op { CUTLASS_HOST_DEVICE diff --git a/include/cutlass/gemm/collective/builders/sm107_blockscaled_umma_builder.inl b/include/cutlass/gemm/collective/builders/sm107_blockscaled_umma_builder.inl new file mode 100644 index 0000000000..ca2e1e425d --- /dev/null +++ b/include/cutlass/gemm/collective/builders/sm107_blockscaled_umma_builder.inl @@ -0,0 +1,256 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ +#pragma once + +#include "cutlass/gemm/collective/builders/sm107_common.inl" +#include "cutlass/gemm/collective/builders/sm100_pipeline_carveout.inl" +#include "cutlass/gemm/collective/builders/sm100_blockscaled_umma_builder.inl" + +///////////////////////////////////////////////////////////////////////////////////////////////// + +namespace cutlass::gemm::collective { + +///////////////////////////////////////////////////////////////////////////////////////////////// + +template < + class ArchTag, + class ElementPairA, + class GmemLayoutATag, + int AlignmentA, + class ElementPairB, + class GmemLayoutBTag, + int AlignmentB, + class ElementAccumulator, + class TileShape_MNK, // (MmaAtomShapeM, MmaAtomShapeN, TileK) + class ClusterShape_MNK, // Static cluster shape or dynamic (int, int, _1) + class StageCountType, + class BuilderScheduleTag +> +struct CollectiveBuilder< + ArchTag, + arch::OpClassBlockScaledTensorOp, + ElementPairA, + GmemLayoutATag, + AlignmentA, + ElementPairB, + GmemLayoutBTag, + AlignmentB, + ElementAccumulator, + TileShape_MNK, + ClusterShape_MNK, + StageCountType, + BuilderScheduleTag, + cute::enable_if_t< + cute::is_same_v + && + // Blockscaled Gemm + ( + cute::is_base_of_v || + cute::is_base_of_v + ) + && + // Alignment check + detail::sm1xx_blockscaled_gemm_is_aligned::data_type, + AlignmentA, + typename detail::blockscaled::blockscaled_type::data_type, + AlignmentB, + BuilderScheduleTag>()>> +{ + using ElementSFA = typename detail::blockscaled::blockscaled_type::sf_type; + using ElementSFB = typename detail::blockscaled::blockscaled_type::sf_type; + using ElementA = typename detail::blockscaled::blockscaled_type::data_type; + using ElementB = typename detail::blockscaled::blockscaled_type::data_type; + using ElementSF = ElementSFA; + + static constexpr cute::UMMA::Major UmmaMajorA = cutlass::gemm::collective::detail::tag_to_umma_major_A(); + static constexpr cute::UMMA::Major UmmaMajorB = cutlass::gemm::collective::detail::tag_to_umma_major_B(); + + static_assert(cute::is_static_v, "TileShape has to be static"); + static_assert(detail::blockscaled::check_input_datatypes(), "Incorrect input types"); + + static constexpr bool is_2sm = detail::blockscaled::is_2sm(); + static constexpr auto Instr = detail::blockscaled::select_instr(); + + static constexpr bool WithBreuse = + (cute::is_base_of_v || + cute::is_base_of_v || + cute::is_base_of_v) + ? true : false; + + static_assert( + cute::is_base_of_v || + cute::is_base_of_v, + "SM107 blockscaled builder only supports MXF8F6F4 (SF vector size 32) and MXNVF4 (SF vector size 16 or 32) kernel schedules."); + static constexpr uint32_t SFVecSize = + cute::is_base_of_v ? 16 : 32; + + using TiledMma = decltype(cutlass::gemm::collective::detail::sm107_make_blockscaled_trivial_tiled_mma< + ElementPairA, ElementPairB, ElementAccumulator, + TileShape_MNK, ClusterShape_MNK, + UmmaMajorA, UmmaMajorB, Instr, BuilderScheduleTag, SFVecSize, WithBreuse>()); + + static constexpr bool UseMxf8f6f4 = Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8; + + static_assert(UseMxf8f6f4 || (cutlass::gemm::detail::is_k_major_A() && cutlass::gemm::detail::is_k_major_B()), "Only MMA.MXF8F6F4 supports non-K major inputs"); + + // Data type used by MMA instruction + using ElementAMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element()); + using ElementBMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element()); + + // Basic storage block for new Scaling Factor Layouts + using AtomThrID = typename TiledMma::AtomThrID; + using Sm107BlkScaledConfig = cutlass::detail::Sm107BlockScaledConfig; + + using ElementAMma_SmemAllocType = cute::conditional_t; + using ElementBMma_SmemAllocType = cute::conditional_t; + + // ((MMA_TILE_M,MMA_TILE_K), MMA_M, MMA_K) + using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(cute::size<0>(TileShape_MNK{}), + cute::size<2>(TileShape_MNK{})))); + // ((MMA_TILE_N,MMA_TILE_K), MMA_N, MMA_K) + using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(cute::size<1>(TileShape_MNK{}), + cute::size<2>(TileShape_MNK{})))); + + using GmemTiledCopyA = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_A( + ClusterShape_MNK{}, AtomThrID{})); + + using GmemTiledCopyB = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_B( + ClusterShape_MNK{}, AtomThrID{})); + + using GmemTiledCopySFA = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_A( + ClusterShape_MNK{}, AtomThrID{})); + + using GmemTiledCopySFB = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_SFB( + ClusterShape_MNK{}, AtomThrID{})); + + using GmemTiledCopyPairA = decltype(cute::make_tuple(GmemTiledCopyA{}, GmemTiledCopySFA{})); + using GmemTiledCopyPairB = decltype(cute::make_tuple(GmemTiledCopyB{}, GmemTiledCopySFB{})); + + // + // Construct SMEM layout (SmemLayoutAtom) for A and SFA + // + using BlockTileA_M = decltype(cute::size<0,0>(MmaShapeA_MK{}) * cute::size<1>(MmaShapeA_MK{})); + using BlockTileA_K = decltype(cute::size<0,1>(MmaShapeA_MK{}) * cute::size<2>(MmaShapeA_MK{})); + using SmemLayoutAtomA = decltype(cutlass::gemm::collective::detail::sm100_smem_selector< + UmmaMajorA, ElementAMma_SmemAllocType, BlockTileA_M, BlockTileA_K>()); + + // A single indivisible block will hold 4 scale factors of 128 rows/columns (A/B matrix). + // 4 is chosen to make consecutive 32bits of data to have scale factors for only a single row (col). 32bits corresponds to the TMEM word size + using Blk_MN = typename Sm107BlkScaledConfig::Blk_MN; + using Blk_SF = typename Sm107BlkScaledConfig::Blk_SF; + using Blk_Elems = decltype(Blk_MN{} * Blk_SF{}); + using SmemLayoutAtomSFA = decltype(Sm107BlkScaledConfig::deduce_smem_layoutSFA(TiledMma{}, TileShape_MNK{})); + using SmemLayoutAtomsA = decltype(cute::make_tuple(SmemLayoutAtomA{}, SmemLayoutAtomSFA{})); + + // + // Construct SMEM layout (SmemLayoutAtom) for B and SFB + // + using BlockTileB_N = decltype(cute::size<0,0>(MmaShapeB_NK{}) * cute::size<1>(MmaShapeB_NK{})); + using BlockTileB_K = decltype(cute::size<0,1>(MmaShapeB_NK{}) * cute::size<2>(MmaShapeB_NK{})); + using SmemLayoutAtomB = decltype(cutlass::gemm::collective::detail::sm100_smem_selector< + UmmaMajorB, ElementBMma_SmemAllocType, BlockTileB_N, BlockTileB_K>()); + using SmemLayoutAtomSFB = decltype(Sm107BlkScaledConfig::deduce_smem_layoutSFB(TiledMma{}, TileShape_MNK{})); + using SmemLayoutAtomsB = decltype(cute::make_tuple(SmemLayoutAtomB{}, SmemLayoutAtomSFB{})); + + // + // Construct Strides for A, SFA, B, and SFB + // + using StrideA = cutlass::gemm::TagToStrideA_t; + using StrideB = cutlass::gemm::TagToStrideB_t; + using InternalStrideA = cute::remove_pointer_t; + using InternalStrideB = cute::remove_pointer_t; + using InternalLayoutSFA = decltype(Sm107BlkScaledConfig::deduce_layoutSFA()); + using InternalLayoutSFB = decltype(Sm107BlkScaledConfig::deduce_layoutSFB()); + using LayoutSFA = cute::conditional_t, InternalLayoutSFA, InternalLayoutSFA *>; + using LayoutSFB = cute::conditional_t, InternalLayoutSFB, InternalLayoutSFB *>; + using StridePairA = decltype(cute::make_tuple(StrideA{}, LayoutSFA{})); + using StridePairB = decltype(cute::make_tuple(StrideB{}, LayoutSFB{})); + + static constexpr int MMA_N = cute::size<1>(TileShape_MNK{}); + static constexpr uint32_t AccumulatorPipelineStageCount = (WithBreuse && (MMA_N == 192 || MMA_N == 256)) ? 1 : 2; + static constexpr bool IsArrayOfPointersGemm = false; + // Grouped GEMM(where Stride type is Stride*) uses specific static tile scheduler. + static constexpr bool IsGroupGemm = false; + static constexpr bool IsRCGroupGemm = false; + static constexpr uint32_t SchedulerPipelineStageCount = cute::conditional_return(8, 2); + + static constexpr uint32_t KernelSmemCarveout = detail::Sm100DenseGemmTmaUmmaCarveout< + ClusterShape_MNK, + AccumulatorPipelineStageCount, + SchedulerPipelineStageCount, + detail::CLCResponseSize, + IsArrayOfPointersGemm, + 4 // 4 Tensor maps for A, SFA, B and SFB + >::KernelSmemCarveout; + // Reduce SMEM capacity available for buffers considering barrier allocations. + + static constexpr int ReducedSmemCapacityBytes = detail::sm100_reduced_smem_capacity_bytes(); + + using SmemTileShape = cute::Shape; + + static constexpr int PipelineStages = cutlass::gemm::collective::detail::sm100_compute_stage_count_or_override_blockscaled< + ReducedSmemCapacityBytes, ElementAMma_SmemAllocType, ElementBMma_SmemAllocType, SmemTileShape, SmemLayoutAtomSFA, SmemLayoutAtomSFB>(StageCountType{}); + static_assert(PipelineStages > 0, "Smem usage is too high. Can't create any SMEM buffers for A, B, SFA, and SFB."); + + using DispatchPolicy = + cutlass::gemm::MainloopSm107TmaUmmaWarpSpecializedBlockScaled< + PipelineStages, + SchedulerPipelineStageCount, + AccumulatorPipelineStageCount, + ClusterShape_MNK, + WithBreuse, + ArchTag + >; + + using CollectiveOp = cutlass::gemm::collective::CollectiveMma< + DispatchPolicy, + TileShape_MNK, + cute::tuple, + StridePairA, + cute::tuple, + StridePairB, + TiledMma, + GmemTiledCopyPairA, + SmemLayoutAtomsA, + void, + cute::identity, + GmemTiledCopyPairB, + SmemLayoutAtomsB, + void, + cute::identity + >; +}; + +///////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace cutlass::gemm::collective + +///////////////////////////////////////////////////////////////////////////////////////////////// diff --git a/include/cutlass/gemm/collective/builders/sm107_common.inl b/include/cutlass/gemm/collective/builders/sm107_common.inl new file mode 100644 index 0000000000..4aeba724e6 --- /dev/null +++ b/include/cutlass/gemm/collective/builders/sm107_common.inl @@ -0,0 +1,359 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#pragma once + +#include "cutlass/gemm/collective/builders/sm100_common.inl" + +///////////////////////////////////////////////////////////////////////////////////////////////// + +namespace cutlass::gemm::collective { + +///////////////////////////////////////////////////////////////////////////////////////////////// + +namespace detail { + +template< + class ElementAMma, + class ElementBMma, + class ElementAMmaccumulator, + class TileShape_MNK, + class ClusterShape_MNK, + UMMA::Major UmmaMajorA, + UMMA::Major UmmaMajorB, + bool WithBreuse, + UMMA::ScaleIn ANeg = UMMA::ScaleIn::One, + UMMA::ScaleIn BNeg = UMMA::ScaleIn::One +> +constexpr auto +sm107_make_1sm_trivial_tiled_mma() { + + constexpr int M = cute::size<0>(TileShape_MNK{}) / (WithBreuse ? 2 : 1); + static_assert(M == 128, "Invalid TileShape_M."); + + constexpr int N = cute::size<1>(TileShape_MNK{}); + static_assert(N % 8 == 0 && N <= 256, "Invalid TileShape_N."); + + if constexpr (cute::is_same_v || + cute::is_same_v || + cute::is_same_v) { + + return make_tiled_mma( + cute::SM107_MMA_F8F6F4_SS< + ElementAMma, + ElementBMma, + ElementAMmaccumulator, + M, + N, + UmmaMajorA, + UmmaMajorB, + ANeg, + BNeg>{} + ); + } + else { + static_assert(cutlass::detail::dependent_false, + "Unsupported configuration for SM107 collective builder."); + } +} + +template< + class ElementAMma, + class ElementBMma, + class ElementAMmaccumulator, + class TileShape_MNK, + class ClusterShape_MNK, + UMMA::Major UmmaMajorA, + UMMA::Major UmmaMajorB, + bool WithBreuse, + UMMA::ScaleIn ANeg = UMMA::ScaleIn::One, + UMMA::ScaleIn BNeg = UMMA::ScaleIn::One +> +constexpr auto +sm107_make_2sm_trivial_tiled_mma() { + + constexpr int M = cute::size<0>(TileShape_MNK{}) / (WithBreuse ? 2 : 1); + static_assert(M == 256, "Invalid TileShape_M."); + + constexpr int N = cute::size<1>(TileShape_MNK{}); + static_assert(N % 8 == 0 && N <= 256, "Invalid TileShape_N."); + + constexpr int K = cute::size<2>(TileShape_MNK{}); + static_assert(K == 128 || K == 64, "Invalid TileShape_K"); + if constexpr (cute::is_same_v || + cute::is_same_v || + cute::is_same_v) { + + // For the case of b-reuse and .2CTA instructions, we permute the M mode + // such that each CTA within a pair contains a consecutive portion of the A tensor + if constexpr (WithBreuse) { + auto permutation_mnk = make_tile( + Layout, Stride<_1, _256, _128>>{}, + cute::Int{}, + cute::Int{}); + return make_tiled_mma( + cute::SM107_MMA_F8F6F4_2x1SM_SS< + ElementAMma, + ElementBMma, + ElementAMmaccumulator, + M, + N, + UmmaMajorA, + UmmaMajorB, + ANeg, + BNeg>{}, + Layout>{}, + permutation_mnk + ); + } else { + return make_tiled_mma( + cute::SM107_MMA_F8F6F4_2x1SM_SS< + ElementAMma, + ElementBMma, + ElementAMmaccumulator, + M, + N, + UmmaMajorA, + UmmaMajorB, + ANeg, + BNeg>{} + ); + } + } + else { + static_assert(cutlass::detail::dependent_false, + "Unsupported configuration for SM107 collective builder."); + } +} + +// For new MMA construction and partitioning that supports both dynamic and static cluster shape. +// Used in conjunction with make_tma_atom_(A|B)_sm100 +// ClusterShape_MNK can be dynamic or static. +template< + class ElementAMma, + class ElementBMma, + class ElementAccumulator, + class TileShape_MNK, + class ClusterShape_MNK, + UMMA::Major UmmaMajorA, + UMMA::Major UmmaMajorB, + class BuilderScheduleTag, + bool WithBreuse = false, + UMMA::ScaleIn ANeg = UMMA::ScaleIn::One, + UMMA::ScaleIn BNeg = UMMA::ScaleIn::One +> +constexpr auto +sm107_make_trivial_tiled_mma() { + // MMA_2SM requested + if constexpr (cute::is_base_of_v) { + return sm107_make_2sm_trivial_tiled_mma< + ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK, + ClusterShape_MNK, UmmaMajorA, UmmaMajorB, WithBreuse, ANeg, BNeg>(); + } + // MMA_1SM requested + else + if constexpr (cute::is_base_of_v) { + return sm107_make_1sm_trivial_tiled_mma< + ElementAMma, ElementBMma, ElementAccumulator, TileShape_MNK, + ClusterShape_MNK, UmmaMajorA, UmmaMajorB, WithBreuse, ANeg, BNeg>(); + } else { + static_assert(cutlass::detail::dependent_false, + "Unsupported configuration for SM107 collective builder."); + } +} + +template < + class ElementPairA, + class ElementPairB, + class ElementAccumulator, + class TileShape_MNK, + class ClusterShape_MNK, + UMMA::Major UmmaMajorA, + UMMA::Major UmmaMajorB, + detail::blockscaled::BlockScaledInstr Instr, + class BuilderScheduleTag, + int SFVecSize, + bool WithBreuse +> +constexpr auto +sm107_make_blockscaled_1sm_trivial_tiled_mma() { + using AtomLayout_MNK = Layout; + constexpr int M = cute::size<0>(TileShape_MNK{}) / (WithBreuse ? 2 : 1); + static_assert(M == 128, "Invalid TileShape_M."); + + constexpr int N = cute::size<1>(TileShape_MNK{}); + static_assert(N == 64 || N == 128 || N == 192 || N == 256, "Invalid TileShape_N."); + + using ElementSFA = typename detail::blockscaled::blockscaled_type::sf_type; + using ElementSFB = typename detail::blockscaled::blockscaled_type::sf_type; + using ElementA = typename detail::blockscaled::blockscaled_type::data_type; + using ElementB = typename detail::blockscaled::blockscaled_type::data_type; + + using ElementAMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element()); + using ElementBMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element()); + + using ElementSF = ElementSFA; + if constexpr (Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8) { + return make_tiled_mma(cute::SM107_MMA_MXF8F6F4_SS{}); + } + else if constexpr (Instr == detail::blockscaled::BlockScaledInstr::MXF4_NVF4) { + return make_tiled_mma(cute::SM107_MMA_MXF4NVF4_SS{}); + } + else { + static_assert(cutlass::detail::dependent_false, + "Unsupported configuration for SM107 tiled MMA."); + } +} + +template < + class ElementPairA, + class ElementPairB, + class ElementAccumulator, + class TileShape_MNK, + class ClusterShape_MNK, + UMMA::Major UmmaMajorA, + UMMA::Major UmmaMajorB, + detail::blockscaled::BlockScaledInstr Instr, + class BuilderScheduleTag, + int SFVecSize, + bool WithBreuse +> +constexpr auto +sm107_make_blockscaled_2sm_trivial_tiled_mma() { + constexpr int M = cute::size<0>(TileShape_MNK{}) / (WithBreuse ? 2 : 1); + static_assert(M == 256, "Invalid TileShape_M."); + + constexpr int N = cute::size<1>(TileShape_MNK{}); + static_assert(N == 64 || N == 128 || N == 192 || N == 256, "Invalid TileShape_N."); + + constexpr int K = cute::size<2>(TileShape_MNK{}); + static_assert((Instr == detail::blockscaled::BlockScaledInstr::MXF4_NVF4) ? (K == 256) : + (Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8) ? (K == 128) : + false, + "Invalid TileShape_K"); + + using ElementSFA = typename detail::blockscaled::blockscaled_type::sf_type; + using ElementSFB = typename detail::blockscaled::blockscaled_type::sf_type; + using ElementA = typename detail::blockscaled::blockscaled_type::data_type; + using ElementB = typename detail::blockscaled::blockscaled_type::data_type; + + using ElementAMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element()); + using ElementBMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element()); + + using ElementSF = ElementSFA; + if constexpr (Instr == detail::blockscaled::BlockScaledInstr::MXF4F6F8) { + // For the case of b-reuse and .2CTA instructions, we permute the M mode + // such that each CTA within a pair contains a consecutive portion of the A tensor + if constexpr (WithBreuse) { + auto permutation_mnk = make_tile( + Layout, Stride<_1, _256, _128>>{}, + cute::Int{}, + cute::Int{}); + return make_tiled_mma( + cute::SM107_MMA_MXF8F6F4_2x1SM_SS{}, + Layout>{}, + permutation_mnk + ); + } + else { + return make_tiled_mma(cute::SM107_MMA_MXF8F6F4_2x1SM_SS{}); + } + } + else if constexpr (Instr == detail::blockscaled::BlockScaledInstr::MXF4_NVF4) { + // For the case of b-reuse and .2CTA instructions, we permute the M mode + // such that each CTA within a pair contains a consecutive portion of the A tensor + if constexpr (WithBreuse) { + auto permutation_mnk = make_tile( + Layout, Stride<_1, _256, _128>>{}, + cute::Int{}, + cute::Int{}); + return make_tiled_mma( + cute::SM107_MMA_MXF4NVF4_2x1SM_SS{}, + Layout>{}, + permutation_mnk + ); + } + else { + return make_tiled_mma(cute::SM107_MMA_MXF4NVF4_2x1SM_SS{}); + } + } + else { + static_assert(cutlass::detail::dependent_false, + "Unsupported configuration for SM107 tiled MMA."); + } +} + +template < + class ElementPairA, + class ElementPairB, + class ElementAccumulator, + class TileShape_MNK, + class ClusterShape_MNK, + UMMA::Major UmmaMajorA, + UMMA::Major UmmaMajorB, + detail::blockscaled::BlockScaledInstr Instr, + class BuilderScheduleTag, + int SFVecSize, + bool WithBreuse = false +> +constexpr auto +sm107_make_blockscaled_trivial_tiled_mma() { + // MMA_2SM requested + if constexpr (cute::is_base_of_v) { + return sm107_make_blockscaled_2sm_trivial_tiled_mma< + ElementPairA, ElementPairB, ElementAccumulator, TileShape_MNK, + ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Instr, BuilderScheduleTag, + SFVecSize, WithBreuse>(); + } + // MMA_1SM requested + else if constexpr (cute::is_base_of_v) { + return sm107_make_blockscaled_1sm_trivial_tiled_mma< + ElementPairA, ElementPairB, ElementAccumulator, TileShape_MNK, + ClusterShape_MNK, UmmaMajorA, UmmaMajorB, Instr, BuilderScheduleTag, + SFVecSize, WithBreuse>(); + } + else { + static_assert(cutlass::detail::dependent_false, + "Unsupported configuration for SM107 tiled mma builder"); + } +} + +} // namespace detail + +///////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace cutlass::gemm::collective diff --git a/include/cutlass/gemm/collective/builders/sm107_sparse_config.inl b/include/cutlass/gemm/collective/builders/sm107_sparse_config.inl new file mode 100644 index 0000000000..20c61505fb --- /dev/null +++ b/include/cutlass/gemm/collective/builders/sm107_sparse_config.inl @@ -0,0 +1,171 @@ +/*************************************************************************************************** + * Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#pragma once + +#include "cute/config.hpp" // CUTE_STATIC_ASSERT +#include "cute/layout.hpp" // cute::Layout, cute::Shape, cute::Stride +#include "cute/numeric/integral_constant.hpp" // cute::Int +#include "cute/numeric/numeric_types.hpp" // cute::sizeof_bits_v +#include "cute/pointer_sparse.hpp" // cute::is_sparse +#include "cute/util/type_traits.hpp" // cute::is_same_v, cute::conditional_t +#include "cutlass/fast_math.h" // cutlass::round_up +#include "cutlass/layout/matrix.h" // cutlass::layout::RowMajor + +namespace cutlass { + +using namespace cute; + +/// Sparse layout configuration for SM107 2:4 FP4 compression. +/// +/// This type prepares A/E tensors for an IR-based SM107 sparse GEMM. It is not +/// a C++ sparse GEMM kernel configuration. +template< + class ElementAMma_, + class LayoutATag_, + class ElementEMma_ +> +struct Sm107GemmSparseConfig { + /// ElementAMma Check + static_assert(cute::is_same_v, "LayoutATag MUST be RowMajor"); // not used, for API compatibility + static_assert(cute::is_sparse::value, "ElementAMma MUST be sparse elem"); + static_assert(cute::is_sparse::value, "ElementEMma MUST be sparse elem"); + + /// A + using ElementAMma = ElementAMma_; + using ElementAMmaRaw = typename ElementAMma::raw_type; // fp4 + using ElementAMmaSparsity = Int; // 2 + + /// MetaData (E) + using ElementEMma = ElementEMma_; + using ElementEMmaRaw = typename ElementEMma::raw_type; // uint8_t + using ElementEMmaSparsity = Int; // 8 + + /// Number of ElementARaw stored in ElementAMmaRaw + using ElemsARawPerElementAMmaRaw = _1; + + /// ElementA Sparsity Ratio + using ElementASparsity = _2; + + // Logical/Physical ElementA per Chunk + // 2:4 sparse for all Rubin + using LogicalElemsAPerChunk = _4; + using PhysicalElemsAPerChunk = Int; + + /// Metadata Bits + using ElementEBitsPerChunk = _4; + using ElementEBitsPerElementAMma = _2; + + /// Metadata Layout + using TensorEAtom = Layout, + Stride<_128, _1>>; + + // Logical elems that construct the atomK for tensorE/A. + using TensorEAtomK = Int(TensorEAtom{})>; + using TensorEAtomM = Int(TensorEAtom{})>; + + using TensorEAlignmentM = TensorEAtomM; + using TensorEAlignmentK = TensorEAtomK; + + // When A is K major, TensorAAlignmentK needs to be multiplier of TMA requirements times tensorA sparsity + // this is b.c. TensorACompressed needs to satisfy TMA requirements. + // (LogicalElemsAPerChunk is always smaller than TMA in this case.) + // NOTE: TensorAAlignmentK already contains the 2x sparsity factor when k-major + using TensorAAlignmentK = Int<128 / cute::sizeof_bits_v>; + using TensorAAlignmentM = _1; // When A is K Major, no requirements on TensorAAlignmentM. + + // For compressor kernel compatibility + static constexpr bool IsTF32 = false; + + // The following two functions are provided for user fill dynamic problem size to the layout_a/e. + template < + class ProblemShape + > + CUTE_HOST_DEVICE + static constexpr auto + fill_layoutA(ProblemShape problem_shape) { + // * Purpose of this function + // This function is sparse gemm equivalent of + // + // ```cpp + // using LayoutATag = cutlass::layout::RowMajor; + // using StrideA = cutlass::gemm::TagToStrideA_t; // ( cute::Stride, int64_t> ) + // auto stride_a = cutlass::make_cute_packed_stride(StrideA{}, cute::make_shape(M, K, L)); // (M, cute::Int<1>, L) + // auto layout_a = cute::make_layout(cute::make_shape(M, K, L), stride_a); + // ``` + // + // Unlike dense gemm where we can simply call `TagToStrideA_t` resp. `make_cute_packed_stride` + // to get the shape and stride, sparse gemm needs to consider the cute::sparse_elem<> representation. + // Thus, it's easier to construct the layout directly. + // + // * NOTE + // 1. Returned layout should be used with `cute::sparse_elem<>` pointer, instead of raw element A ptr + // 2. `TensorAAlignmentK` already include 2x sparsity factor along K dim. + const auto [M, N, K, L] = problem_shape; + + // Round up to satisfy TensorA Alignment requirement + const auto M_AlignedA = cutlass::round_up(M, TensorAAlignmentM{}); + const auto K_AlignedA = cutlass::round_up(K, TensorAAlignmentK{}); + + return make_layout( + make_shape(int32_t(M_AlignedA), + make_shape(ElementASparsity{}, int32_t(K_AlignedA / ElementASparsity{})), + int32_t(L)), + make_stride(int64_t(K_AlignedA), + make_stride(_1{}, ElementASparsity{}), + (L == 1) ? int64_t(0) : int64_t(M_AlignedA * K_AlignedA)) + ); + } + + template < + class ProblemShape + > + CUTE_HOST_DEVICE + static constexpr auto + fill_layoutE(ProblemShape problem_shape) { + const auto [M, N, K, L] = problem_shape; + + // Round up to satisfy TensorEAlignment requirement + const auto M_AlignedE = cutlass::round_up(M, TensorEAlignmentM{}); + const auto K_AlignedE = cutlass::round_up(K, TensorEAlignmentK{}); + + return make_layout( + make_shape(make_shape(shape<0>(TensorEAtom{}), int32_t(M_AlignedE / TensorEAlignmentM{})), + make_shape(shape<1>(TensorEAtom{}), int32_t(K_AlignedE / TensorEAlignmentK{})), + int32_t(L)), + make_stride(make_stride(stride<0>(TensorEAtom{}), cute::Int{}), + make_stride(stride<1>(TensorEAtom{}), int64_t(M_AlignedE * TensorEAlignmentK{})), + (L == 1) ? int64_t(0) : int64_t(M_AlignedE * K_AlignedE)) + ); + } +}; + +} // namespace cutlass diff --git a/include/cutlass/gemm/collective/builders/sm107_umma_builder.inl b/include/cutlass/gemm/collective/builders/sm107_umma_builder.inl new file mode 100644 index 0000000000..3eb607be88 --- /dev/null +++ b/include/cutlass/gemm/collective/builders/sm107_umma_builder.inl @@ -0,0 +1,259 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ +#pragma once + +#include "cutlass/gemm/collective/builders/sm107_common.inl" +#include "cutlass/gemm/collective/builders/sm100_pipeline_carveout.inl" +#include "cutlass/gemm/collective/builders/sm100_umma_builder.inl" + +///////////////////////////////////////////////////////////////////////////////////////////////// + +namespace cutlass::gemm::collective { + +///////////////////////////////////////////////////////////////////////////////////////////////// + +namespace detail { + +template +CUTLASS_HOST_DEVICE +static constexpr bool +sm107_check_input_datatypes() { + auto is_f4f6f8_input = [&]() { + // Allowed input element datatype for narrow precision GEMM + return ( + ( + cute::is_same_v || + cute::is_same_v || + cute::is_same_v + ) && + ( + cute::is_same_v || + cute::is_same_v || + cute::is_same_v + ) + ) || + ( + ( + cute::is_same_v || + cute::is_same_v || + cute::is_same_v || + cute::is_same_v || + cute::is_same_v + ) && + ( + cute::is_same_v || + cute::is_same_v || + cute::is_same_v || + cute::is_same_v || + cute::is_same_v + ) + ); + }; + + static_assert(is_f4f6f8_input(), "Unsupported data type for ElementA"); + + return true; +} + +} // namespace detail + +///////////////////////////////////////////////////////////////////////////////////////////////// + +template < + class ArchTag, + class ElementA, + class GmemLayoutATag, + int AlignmentA, + class ElementB, + class GmemLayoutBTag, + int AlignmentB, + class ElementAccumulator, + class TileShape_MNK, + class ClusterShape_MNK, + class StageCountType, + class BuilderScheduleTag +> +struct CollectiveBuilder< + ArchTag, + arch::OpClassTensorOp, + ElementA, + GmemLayoutATag, + AlignmentA, + ElementB, + GmemLayoutBTag, + AlignmentB, + ElementAccumulator, + TileShape_MNK, // (MmaAtomShapeM, MmaAtomShapeN, TileK) + ClusterShape_MNK, // Static cluster shape or dynamic (int, int, _1) + StageCountType, + BuilderScheduleTag, + cute::enable_if_t< + cute::is_same_v && + not cute::is_tuple_v && not cute::is_tuple_v && + // Dense Gemm + ( + cute::is_base_of_v || + cute::is_base_of_v + ) && + // Alignment check + detail::sm1xx_gemm_is_aligned()>> +{ + static_assert(cute::is_static_v, "TileShape has to be static"); + static_assert(detail::sm107_check_input_datatypes(), "Incorrect input types"); + + static constexpr cute::UMMA::Major UmmaMajorA = cutlass::gemm::collective::detail::tag_to_umma_major_A(); + static constexpr cute::UMMA::Major UmmaMajorB = cutlass::gemm::collective::detail::tag_to_umma_major_B(); + + // Data type used by MMA instruction + using ElementAMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element()); + using ElementBMma = decltype(cutlass::gemm::collective::detail::sm1xx_kernel_input_element_to_mma_input_element()); + + static constexpr bool is_2sm = cute::is_base_of_v || + (not cute::is_base_of_v && + not cute::is_base_of_v && + cute::is_static_v && + cute::get<0>(ClusterShape_MNK{}) % 2 == 0 ); + + + static constexpr bool WithBreuse = + cute::is_base_of_v + ? true : false; + + using TiledMma = decltype(detail::sm107_make_trivial_tiled_mma< + ElementAMma, ElementBMma, ElementAccumulator, + decltype(cute::product_each(TileShape_MNK{})), ClusterShape_MNK, + UmmaMajorA, UmmaMajorB, BuilderScheduleTag, WithBreuse>()); + + using ElementAMma_SmemAllocType = cute::conditional_t < 8, uint8_t, ElementAMma>; + using ElementBMma_SmemAllocType = cute::conditional_t < 8, uint8_t, ElementBMma>; + + using AtomThrID = typename TiledMma::AtomThrID; + + using AtomThrShapeMNK = cute::Shape(typename TiledMma::ThrLayoutVMNK{})), _1, _1>; + using CtaTileShape_MNK = decltype(cute::shape_div(TileShape_MNK{}, AtomThrShapeMNK{})); + + // ((MMA_TILE_M,MMA_TILE_K), MMA_M, MMA_K) + using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(cute::size<0>(TileShape_MNK{}), + cute::size<2>(TileShape_MNK{})))); + // ((MMA_TILE_N,MMA_TILE_K), MMA_N, MMA_K) + using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(cute::size<1>(TileShape_MNK{}), + cute::size<2>(TileShape_MNK{})))); + + using BlockTileA_M = decltype(cute::size<0,0>(MmaShapeA_MK{}) * cute::size<1>(MmaShapeA_MK{})); + using BlockTileA_K = decltype(cute::size<0,1>(MmaShapeA_MK{}) * cute::size<2>(MmaShapeA_MK{})); + using BlockTileB_N = decltype(cute::size<0,0>(MmaShapeB_NK{}) * cute::size<1>(MmaShapeB_NK{})); + using BlockTileB_K = decltype(cute::size<0,1>(MmaShapeB_NK{}) * cute::size<2>(MmaShapeB_NK{})); + + // Right divide TileShape_M/N by 1SM/2SM + // Future work: fix partition_shape to account for hierarchies and + // contiguity so we can pass BlockTileA/B to sm100_smem_selector instead + using SmemShape_M = decltype(shape_div(shape<0>(TileShape_MNK{}), shape_div(shape<0>(TileShape_MNK{}), size<0>(TileShape_MNK{}) / size(AtomThrID{})))); + using SmemShape_N = decltype(shape_div(shape<1>(TileShape_MNK{}), shape_div(shape<1>(TileShape_MNK{}), size<1>(TileShape_MNK{}) / size(AtomThrID{})))); + using SmemShape_K = decltype(cute::get<2>(TileShape_MNK{})); + + using GmemTiledCopyA = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_A( + ClusterShape_MNK{}, AtomThrID{})); + using GmemTiledCopyB = decltype(cutlass::gemm::collective::detail::sm100_cluster_shape_to_tma_atom_B( + ClusterShape_MNK{}, AtomThrID{})); + + using SmemLayoutAtomA = decltype(cutlass::gemm::collective::detail::sm100_smem_selector< + UmmaMajorA, ElementAMma_SmemAllocType, SmemShape_M, SmemShape_K>()); + using SmemLayoutAtomB = decltype(cutlass::gemm::collective::detail::sm100_smem_selector< + UmmaMajorB, ElementBMma_SmemAllocType, SmemShape_N, SmemShape_K>()); + static constexpr uint32_t TotalTmemRows = 128; + static constexpr uint32_t Sm107TmemCapacityColumns = cutlass::arch::Sm107::kTmemCapacityColumns; + static constexpr uint32_t TotalTmem = TotalTmemRows * Sm107TmemCapacityColumns; + static constexpr uint32_t AccumulatorPipelineStageCount_ = TotalTmem / (cute::size<0>(CtaTileShape_MNK{}) * cute::size<1>(CtaTileShape_MNK{})); + // 4 accumulator stages works well to buffer the accumulators, while also preventing overhead in the epilogue tail on small tile sizes. + static constexpr uint32_t AccumulatorPipelineStageCount = cute::min(4u, AccumulatorPipelineStageCount_); // Cap at 4 accumulator stages + static_assert(AccumulatorPipelineStageCount > 0, "Accumulator pipeline stage count must be positive. This error probably means that TileShape_MNK and/or TiledMma::ThrLayoutVMNK are wrong."); + + // Calculate scheduler pipeline stages. Having one more stage than the accumulator allows more latency hiding. + using StrideA = cutlass::gemm::TagToStrideA_t; + using StrideB = cutlass::gemm::TagToStrideB_t; + using InternalStrideA = cute::remove_pointer_t; + using InternalStrideB = cute::remove_pointer_t; + + static constexpr bool IsArrayOfPointersGemm = false; + static constexpr bool IsGroupGemm = false; + static constexpr bool IsRCGroupGemm = false; + + static constexpr uint32_t SchedulerPipelineStageCount = cute::conditional_return(8, 2); + + static constexpr uint32_t KernelSmemCarveout = detail::Sm100DenseGemmTmaUmmaCarveout< + ClusterShape_MNK, + AccumulatorPipelineStageCount, + SchedulerPipelineStageCount, + detail::CLCResponseSize, + IsArrayOfPointersGemm + >::KernelSmemCarveout; + // Reduce SMEM capacity available for buffers considering barrier allocations. + static constexpr int ReducedSmemCapacityBytes = detail::sm100_reduced_smem_capacity_bytes(); + + using SmemTileShape = cute::Shape; + using MainloopPipelineStorage = typename cutlass::PipelineTmaUmmaAsync<1>::SharedStorage; + + static constexpr int PipelineStages = cutlass::gemm::collective::detail::sm100_compute_stage_count_or_override< + ReducedSmemCapacityBytes, ElementAMma_SmemAllocType, ElementBMma_SmemAllocType, SmemTileShape, MainloopPipelineStorage>(StageCountType{}); + static_assert(PipelineStages > 0, "Smem usage is too high. Can't create any SMEM buffers for A, and B."); + + using DispatchPolicy = + cutlass::gemm::MainloopSm107TmaUmmaWarpSpecialized< + PipelineStages, + SchedulerPipelineStageCount, + AccumulatorPipelineStageCount, + ClusterShape_MNK, + WithBreuse, + ArchTag + >; + + using CollectiveOp = cutlass::gemm::collective::CollectiveMma< + DispatchPolicy, + TileShape_MNK, + ElementA, + cutlass::gemm::TagToStrideA_t, + ElementB, + cutlass::gemm::TagToStrideB_t, + TiledMma, + GmemTiledCopyA, + SmemLayoutAtomA, + void, + cute::identity, + GmemTiledCopyB, + SmemLayoutAtomB, + void, + cute::identity + >; +}; + +} // namespace cutlass::gemm::collective + +///////////////////////////////////////////////////////////////////////////////////////////////// diff --git a/include/cutlass/gemm/collective/builders/sm1xx_common.inl b/include/cutlass/gemm/collective/builders/sm1xx_common.inl index 763274cb1f..5d573d659b 100644 --- a/include/cutlass/gemm/collective/builders/sm1xx_common.inl +++ b/include/cutlass/gemm/collective/builders/sm1xx_common.inl @@ -155,6 +155,7 @@ constexpr uint32_t find_vector_size() { || cute::is_same_v || cute::is_same_v || cute::is_same_v + || cute::is_base_of_v ) { return 16; } @@ -168,7 +169,8 @@ constexpr uint32_t find_vector_size() { cute::is_same_v || cute::is_same_v || cute::is_same_v || - cute::is_same_v) { + cute::is_base_of_v || + cute::is_same_v) { return 32; } else if constexpr (cute::is_same_v || @@ -532,8 +534,9 @@ check_input_datatypes() { static_assert(!is_auto_instr_selection_policy(), "Auto instr selection isn't valid if scale factor vector size can't be determined from the types"); } - static_assert(cute::is_same_v - || cute::is_same_v, "Incorrect scale factor type"); + static_assert(cute::is_same_v + || cute::is_same_v + || cute::is_same_v, "Incorrect scale factor type"); if constexpr (((sizeof_bits_v == 4 || sizeof_bits_v == 6 || sizeof_bits_v == 8) && (sizeof_bits_v == 4 || sizeof_bits_v == 6 || sizeof_bits_v == 8) ) && // A and B are 4, 6, or 8 bit types and @@ -589,6 +592,7 @@ check_input_datatypes() { || (SfVectorSizeA == 64 && cute::is_base_of_v) || (SfVectorSizeA == 32 && cute::is_base_of_v) || (SfVectorSizeA == 64 && cute::is_base_of_v) + || (SfVectorSizeA == 32 && cute::is_base_of_v) ), "Incorrect SfVectorSize for MX_F4F6F8 is deduced."); // 4. Check the kernel policy. Kernel policy should be either auto or *MXf8f6f4* @@ -597,6 +601,7 @@ check_input_datatypes() { || cute::is_base_of_v || cute::is_base_of_v || cute::is_base_of_v + || cute::is_base_of_v || is_auto_instr_selection_policy()), "Incorrect Kernel Schedule Policy for Mx_F4F6F8 type inputs."); return true; @@ -608,9 +613,11 @@ check_input_datatypes() { /////////////////////////////////////////////////////////////////////// // 1. Check Scale factor data type - static_assert(cute::is_same_v + static_assert(cute::is_same_v || cute::is_same_v - , "MXNV_F4 supports ue8m0 and ue4m3 SF types"); + || (cute::is_same_v + && cute::is_base_of_v) + , "MXNV_F4 supports ue8m0 and ue4m3 SF types; ue5m3 is only supported on SM107"); // 2. Check whether A and B type combinations are valid or not static_assert( ( // If runtime datatypes are used, then both A and B should be runtime data type @@ -637,6 +644,7 @@ check_input_datatypes() { cute::is_base_of_v || cute::is_base_of_v || cute::is_base_of_v || + cute::is_base_of_v || is_auto_instr_selection_policy()), "Incorrect Kernel Schedule Policy for F4 type inputs."); // If a policy is specified, do more checks @@ -659,7 +667,8 @@ check_input_datatypes() { || cute::is_base_of_v || cute::is_base_of_v || cute::is_base_of_v - || cute::is_base_of_v) { + || cute::is_base_of_v + || cute::is_base_of_v) { static_assert((UmmaMajorA == UMMA::Major::K && UmmaMajorB == UMMA::Major::K), "MX/NV_F4 only supports RowMajor A, and ColMajorB"); static_assert(detail::find_vector_size() == SfVectorSizeA, "Kernel Schedule policy doesn't match the scale factor vector size."); @@ -744,14 +753,16 @@ select_instr() { || cute::is_base_of_v || cute::is_base_of_v || cute::is_base_of_v - || cute::is_base_of_v) { + || cute::is_base_of_v + || cute::is_base_of_v) { return detail::blockscaled::BlockScaledInstr::MXF4F6F8; } else if constexpr (cute::is_base_of_v || cute::is_base_of_v || cute::is_base_of_v || cute::is_base_of_v - || cute::is_base_of_v) { + || cute::is_base_of_v + || cute::is_base_of_v) { return detail::blockscaled::BlockScaledInstr::MXF4_NVF4; } else { @@ -770,6 +781,7 @@ select_instr() { || (SfVectorSize == 64 && cute::is_base_of_v || (SfVectorSize == 32 && cute::is_base_of_v) || (SfVectorSize == 64 && cute::is_base_of_v) + || (SfVectorSize == 32 && cute::is_base_of_v) ), "Incorrect SfVectorSize for MX_F4F6F8 is deduced."); return detail::blockscaled::BlockScaledInstr::MXF4F6F8; } diff --git a/include/cutlass/gemm/collective/collective_builder.hpp b/include/cutlass/gemm/collective/collective_builder.hpp index 3d96695a3e..bf83950132 100644 --- a/include/cutlass/gemm/collective/collective_builder.hpp +++ b/include/cutlass/gemm/collective/collective_builder.hpp @@ -52,6 +52,8 @@ #include "cutlass/gemm/collective/builders/sm100_mixed_tma_cpasync_umma_builder.inl" #include "cutlass/gemm/collective/builders/sm100_blockscaled_mixed_tma_cpasync_umma_builder.inl" #include "cutlass/gemm/collective/builders/sm103_blockscaled_umma_builder.inl" +#include "cutlass/gemm/collective/builders/sm107_umma_builder.inl" +#include "cutlass/gemm/collective/builders/sm107_blockscaled_umma_builder.inl" #include "cutlass/gemm/collective/builders/sm120_mma_builder.inl" #include "cutlass/gemm/collective/builders/sm120_blockscaled_mma_builder.inl" #include "cutlass/gemm/collective/builders/sm120_sparse_mma_builder.inl" diff --git a/include/cutlass/gemm/collective/collective_mma.hpp b/include/cutlass/gemm/collective/collective_mma.hpp index ce3450b00d..5967a38a70 100644 --- a/include/cutlass/gemm/collective/collective_mma.hpp +++ b/include/cutlass/gemm/collective/collective_mma.hpp @@ -54,6 +54,8 @@ #include "cutlass/gemm/collective/sm90_mma_array_tma_gmma_ss_warpspecialized_fp8_blockwise_scaling.hpp" #if !defined(__CUDACC_RTC__) #include "cutlass/gemm/collective/sm100_mma_warpspecialized.hpp" +#include "cutlass/gemm/collective/sm107_mma_warpspecialized.hpp" +#include "cutlass/gemm/collective/sm107_blockscaled_mma_warpspecialized.hpp" #include "cutlass/gemm/collective/sm100_mma_array_warpspecialized.hpp" #include "cutlass/gemm/collective/sm100_mma_array_warpspecialized_rcggemm.hpp" #include "cutlass/gemm/collective/sm100_blockscaled_mma_array_warpspecialized_rcggemm.hpp" diff --git a/include/cutlass/gemm/collective/sm107_blockscaled_mma_warpspecialized.hpp b/include/cutlass/gemm/collective/sm107_blockscaled_mma_warpspecialized.hpp new file mode 100644 index 0000000000..3b99e4b851 --- /dev/null +++ b/include/cutlass/gemm/collective/sm107_blockscaled_mma_warpspecialized.hpp @@ -0,0 +1,1116 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ +#pragma once + +#include "cutlass/cutlass.h" +#include "cutlass/detail/collective.hpp" +#include "cutlass/detail/cluster.hpp" +#include "cutlass/gemm/dispatch_policy.hpp" +#include "cutlass/numeric_types.h" +#include "cutlass/pipeline/pipeline.hpp" +#include "cutlass/gemm/gemm.h" +#include "cutlass/detail/sm107_blockscaled_layout.hpp" +#include "cutlass/trace.h" +#include "cutlass/kernel_hardware_info.hpp" +#include "cutlass/detail/collective.hpp" +#include "cutlass/detail/sm100_tmem_helper.hpp" + +#include "cute/algorithm/functional.hpp" +#include "cute/arch/cluster_sm90.hpp" +#include "cute/atom/mma_atom.hpp" +#include "cute/algorithm/gemm.hpp" +#include "cute/numeric/arithmetic_tuple.hpp" + +///////////////////////////////////////////////////////////////////////////////////////////////// + +namespace cutlass::gemm::collective { +using namespace cute; + +// A 4x32dp128bit S2T copy replicates one SMEM +// core matrix across BroadcastFactor TMEM core matrices, but make_utccp_copy() builds +// its tiler from the TMEM (destination) side only. Partitioning the SMEM source without +// spelling that replication out as an explicit zero-stride mode on the MN part of mode 0 +// (hierarchical (MN,K)) makes partition_S() produce a Rest-mode smaller than TMEM's -- +// exactly what happens for VS16, whose MN mode's replication slot is otherwise degenerate +// (size 1), since a VS16 scale factor's MMA mode spans two core matrices. +template +CUTLASS_DEVICE auto +sm107_append_s2t_broadcast_mode(Tensor const& smem_tensor) +{ + if constexpr (BroadcastFactor == 1) { + return smem_tensor; + } else { + auto mn = get<0,0>(smem_tensor.layout()); + auto k = get<0,1>(smem_tensor.layout()); + auto mn_bcast = append(mn, make_layout(Int{}, Int<0>{})); + auto mode0_new = make_layout(mn_bcast, k); + auto new_layout = replace<0>(smem_tensor.layout(), mode0_new); + return make_tensor(smem_tensor.data(), new_layout); + } +} + +///////////////////////////////////////////////////////////////////////////////////////////////// + +// WarpSpecialized Mainloop +// Both DMA Load and MMA methods of this class must be run by a single thread that's picked by elect_one +template < + int Stages, + int SchedulerPipelineStageCount, + int AccumulatorPipelineStageCount, + class ClusterShape, // Static cluster shape or dynamic (int, int, _1) + bool WithBreuse_, + class TileShape_, // (MmaAtomShapeM, MmaAtomShapeN, TileK) + class ElementPairA_, + class StridePairA_, + class ElementPairB_, + class StridePairB_, + class TiledMma_, + class GmemTiledCopyPairA_, + class SmemLayoutAtomPairA_, + class SmemCopyAtomA_, + class TransformA_, + class GmemTiledCopyPairB_, + class SmemLayoutAtomPairB_, + class SmemCopyAtomB_, + class TransformB_> +struct CollectiveMma< + MainloopSm107TmaUmmaWarpSpecializedBlockScaled< + Stages, + SchedulerPipelineStageCount, + AccumulatorPipelineStageCount, + ClusterShape, + WithBreuse_, + arch::Sm107>, + TileShape_, + ElementPairA_, + StridePairA_, + ElementPairB_, + StridePairB_, + TiledMma_, + GmemTiledCopyPairA_, + SmemLayoutAtomPairA_, + SmemCopyAtomA_, + TransformA_, + GmemTiledCopyPairB_, + SmemLayoutAtomPairB_, + SmemCopyAtomB_, + TransformB_> +{ + // + // Type Aliases + // + using TiledMma = TiledMma_; + using AtomThrShapeMNK = Shape(typename TiledMma::ThrLayoutVMNK{})), _1, _1>; + + static constexpr bool WithBreuse = WithBreuse_; + + using DispatchPolicy = MainloopSm107TmaUmmaWarpSpecializedBlockScaled< + Stages, + SchedulerPipelineStageCount, + AccumulatorPipelineStageCount, + ClusterShape, + WithBreuse, + arch::Sm107>; + using ArchTag = typename DispatchPolicy::ArchTag; + using TileShape = TileShape_; + using TiledMMA_SF = TiledMMA, + Layout>, + Tile>; + + static constexpr bool IsDynamicCluster = not cute::is_static_v; + static constexpr int SFVecSize = TiledMma::SFVecSize; + + CUTE_STATIC_ASSERT_V(evenly_divides(TileShape{}, tile_shape(TiledMma{})), + "Static cluster shape used: TileShape should be evenly divided by TiledMma"); + + using CtaShape_MNK = decltype(shape_div(TileShape{}, AtomThrShapeMNK{})); + static_assert(shape<1>(CtaShape_MNK{}) == 192 or shape<1>(CtaShape_MNK{}) == 64 or + shape<1>(CtaShape_MNK{}) == 128 or shape<1>(CtaShape_MNK{}) == 256, + "Cta N should be one of 64/128/192/256"); + + using ClusterTileShape = decltype(make_shape(get<0>(TileShape{})*get<0>(ClusterShape{}),get<1>(TileShape{})*get<1>(ClusterShape{}),get<2>(TileShape{})*get<2>(ClusterShape{}))); + using Sm1xxBlkScaledConfig = cutlass::detail::Sm107BlockScaledConfig; + using Blk_MN = typename Sm1xxBlkScaledConfig::Blk_MN; + static constexpr int IsCtaN192 = shape<1>(CtaShape_MNK{}) == 192; + static constexpr int IsCtaN64 = shape<1>(CtaShape_MNK{}) == 64; + static int constexpr CTA_N_SF = cutlass::ceil_div(size<1>(CtaShape_MNK{}), Blk_MN{}) * Blk_MN{}; + // Tile shape used for partitioning Scale Factor B. + // The M-dim does not affect the SFB, so just set it as the original TileShape; + using TileShape_SF = decltype(make_shape(get<0>(CtaShape_MNK{}), + Int{} * shape<2>(typename TiledMma::ThrLayoutVMNK()), + get<2>(TileShape{}))); + + // Define A and B block shapes for reduced size TMA_LOADs + using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(size<0>(TileShape{}), size<2>(TileShape{})))); + using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(size<1>(TileShape{}), size<2>(TileShape{})))); + + using ElementPairA = ElementPairA_; + using ElementPairB = ElementPairB_; + using ElementAMma = typename TiledMma::ValTypeA; + using ElementBMma = typename TiledMma::ValTypeB; + using StridePairA = StridePairA_; + using StridePairB = StridePairB_; + using SmemLayoutAtomPairA = SmemLayoutAtomPairA_; + using SmemLayoutAtomPairB = SmemLayoutAtomPairB_; + static_assert(cute::is_same_v(ElementPairA{}))>, + remove_cvref_t(ElementPairB{}))>>, "SFA and SFB data types should be the same"); + + // A and B matrices + using ElementA = remove_cvref_t(ElementPairA{}))>; + using StrideA = remove_cvref_t(StridePairA{}))>; + + using ElementB = remove_cvref_t(ElementPairB{}))>; + using StrideB = remove_cvref_t(StridePairB{}))>; + + static constexpr bool IsRuntimeDataTypeA = cutlass::gemm::collective::detail::is_sm10x_runtime_f8f6f4(); + static constexpr bool IsRuntimeDataTypeB = cutlass::gemm::collective::detail::is_sm10x_runtime_f8f6f4(); + + static_assert((IsRuntimeDataTypeA && IsRuntimeDataTypeB) || + (!IsRuntimeDataTypeA && !IsRuntimeDataTypeB), + "ElementA and ElementB should be both runtime or both static."); + + static constexpr bool IsRuntimeDataType = IsRuntimeDataTypeA && IsRuntimeDataTypeB; + + // SFA and SFB + using ElementSF = remove_cvref_t(ElementPairA{}))>; + using LayoutSFA = remove_cvref_t(StridePairA{}))>; + using LayoutSFB = remove_cvref_t(StridePairB{}))>; + + using ElementAccumulator = typename TiledMma::ValTypeC; + using GmemTiledCopyPairA = GmemTiledCopyPairA_; + using GmemTiledCopyPairB = GmemTiledCopyPairB_; + using GmemTiledCopyA = remove_cvref_t(GmemTiledCopyPairA{}))>; + using GmemTiledCopySFA = remove_cvref_t(GmemTiledCopyPairA{}))>; + using GmemTiledCopyB = remove_cvref_t(GmemTiledCopyPairB{}))>; + using GmemTiledCopySFB = remove_cvref_t(GmemTiledCopyPairB{}))>; + + using SmemLayoutAtomA = remove_cvref_t(SmemLayoutAtomPairA{}))>; + using SmemLayoutAtomSFA = remove_cvref_t(SmemLayoutAtomPairA{}))>; + using SmemLayoutAtomB = remove_cvref_t(SmemLayoutAtomPairB{}))>; + using SmemLayoutAtomSFB = remove_cvref_t(SmemLayoutAtomPairB{}))>; + + using SmemCopyAtomA = SmemCopyAtomA_; + using SmemCopyAtomB = SmemCopyAtomB_; + using TransformA = TransformA_; + using TransformB = TransformB_; + + using MainloopPipeline = cutlass::PipelineTmaUmmaAsync< + DispatchPolicy::Stages, + ClusterShape, + AtomThrShapeMNK>; + using MainloopPipelineState = typename MainloopPipeline::PipelineState; + + static_assert(rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)"); + static_assert((size<0>(TileShape{}) % size<0>(SmemLayoutAtomA{})) == 0, "SmemLayoutAtomA must evenly divide the tile shape."); + static_assert((size<2>(TileShape{}) % size<1>(SmemLayoutAtomA{})) == 0, "SmemLayoutAtomA must evenly divide the tile shape."); + static_assert(cute::is_void_v, + "SM100 UMMA cannot have a non-void copy atom for smem sourced instructions."); + + static_assert(rank(SmemLayoutAtomB{}) == 2, "SmemLayoutAtom must be rank 2 (M/N, K)"); + static_assert((size<1>(TileShape{}) % size<0>(SmemLayoutAtomB{})) == 0, "SmemLayoutAtomB must evenly divide the tile shape."); + static_assert((size<2>(TileShape{}) % size<1>(SmemLayoutAtomB{})) == 0, "SmemLayoutAtomB must evenly divide the tile shape."); + static_assert(cute::is_void_v, + "SM100 UMMA cannot have a non-void copy atom for smem sourced instructions."); + + // Tile along K mode first before tiling over MN. PIPE mode last as usual. + // This maximizes TMA boxes due to better smem-K vectorization, reducing total issued TMAs. + // (MMA_TILE_M,MMA_TILE_K),MMA_M,MMA_K,PIPE) + using SmemLayoutA = decltype(UMMA::tile_to_mma_shape( + SmemLayoutAtomA{}, + append(MmaShapeA_MK{}, Int{}), + cute::conditional_t(), Step<_2,_1,_3>, Step<_1,_2,_3>>{})); + // (MMA_TILE_N,MMA_TILE_K),MMA_N,MMA_K,PIPE) + using SmemLayoutB = decltype(UMMA::tile_to_mma_shape( + SmemLayoutAtomB{}, + append(MmaShapeB_NK{}, Int{}), + cute::conditional_t(), Step<_2,_1,_3>, Step<_1,_2,_3>>{})); + + // SmemLayoutAtomSFA and SmemLayoutAtomSFB are for whole CTA tiles. We add the number of pipeline stages here. + // The number of pipeline stages is the same as the number of pipeline stages from AB Load <-> MainLoop + using SmemLayoutSFA = decltype(make_layout( + append(shape(SmemLayoutAtomSFA{}), Int{}), + append(stride(SmemLayoutAtomSFA{}), size(filter_zeros(SmemLayoutAtomSFA{}))) + )); + using SmemLayoutSFB = decltype(make_layout( + append(shape(SmemLayoutAtomSFB{}), Int{}), + append(stride(SmemLayoutAtomSFB{}), size(filter_zeros(SmemLayoutAtomSFB{}))) + )); + + static_assert(cute::is_base_of::value && + cute::is_base_of::value, + "MMA atom must source both A and B operand from smem_desc for this mainloop."); + static_assert( + (size(AtomThrShapeMNK{}) == 1 && + (cute::is_same_v || cute::is_same_v)) || + (size(AtomThrShapeMNK{}) == 2 && + (cute::is_same_v || cute::is_same_v)), + "GmemTiledCopy - invalid TMA copy atom specified."); + static_assert( + (size(AtomThrShapeMNK{}) == 1 && + (cute::is_same_v || cute::is_same_v)) || + (size(AtomThrShapeMNK{}) == 2 && + (cute::is_same_v || cute::is_same_v)), + "GmemTiledCopy - invalid TMA copy atom specified."); + + static constexpr bool IsF8F6F4 = detail::is_sm100_mma_f8f6f4(); + + using TmaInternalElementA = cute::conditional_t; + using TmaInternalElementB = cute::conditional_t; + + using SmemAllocTypeA = cute::conditional_t < 8, uint8_t, ElementAMma>; + using SmemAllocTypeB = cute::conditional_t < 8, uint8_t, ElementBMma>; + + using BitTypeElementA = cute::uint_bit_t>; + using BitTypeElementB = cute::uint_bit_t>; + + using ArrayElementA = cute::conditional_t; + using ArrayElementB = cute::conditional_t; + + using RuntimeDataTypeA = typename detail::sm10x_block_scale_runtime_input_t::Type; + using RuntimeDataTypeB = typename detail::sm10x_block_scale_runtime_input_t::Type; + + struct SharedStorage { + struct TensorStorage : cute::aligned_struct<128, _0> { + cute::ArrayEngine> smem_A; + cute::ArrayEngine> smem_B; + cute::ArrayEngine> smem_SFA; + cute::ArrayEngine> smem_SFB; + } tensors; + + using PipelineStorage = typename MainloopPipeline::SharedStorage; + PipelineStorage pipeline; + }; + + // Expose shared storage for tensors/pipelines separately to allow kernel layer to reorder them. + using TensorStorage = typename SharedStorage::TensorStorage; + using PipelineStorage = typename SharedStorage::PipelineStorage; + + // Only one thread issues the TMA and updates the barriers in a 2SM MMA, adjust bytes accordingly + static constexpr uint32_t SFTransactionBytes = + cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutSFA{})) * cute::sizeof_bits_v) + + cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutSFB{})) * cute::sizeof_bits_v); + static constexpr uint32_t ABTmaTransactionBytes = + cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutA{})) * cute::sizeof_bits_v) + + cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutB{})) * cute::sizeof_bits_v); + static constexpr uint32_t TmaTransactionBytes = ABTmaTransactionBytes + SFTransactionBytes; + + template + struct TmemStorage { + AccTensor accumulators; + SfaTensor tCtSFA; + SfbTensor tCtSFB; + }; + + template < + class KTileCount, + class GTensorPartitionedA, class GTensorPartitionedB, + class STensorA, class STensorB, + class GTensorPartitionedSFA, class GTensorPartitionedSFB, + class STensorSFA, class STensorSFB + > + struct LoadParams { + // for scheduler + KTileCount k_tiles; + // for input tensor values + GTensorPartitionedA tAgA_mkl; + GTensorPartitionedB tBgB_nkl; + STensorA tAsA; + STensorB tBsB; + // for scale factor tensor values + GTensorPartitionedSFA tAgSFA_mkl; + GTensorPartitionedSFB tBgSFB_nkl; + STensorSFA tAsSFA; + STensorSFB tBsSFB; + // the TMA multicast masks + uint16_t mcast_mask_a; + uint16_t mcast_mask_b; + uint16_t mcast_mask_sfa; + uint16_t mcast_mask_sfb; + + CUTLASS_DEVICE + LoadParams ( + KTileCount k_tiles_, + GTensorPartitionedA tAgA_mkl_, GTensorPartitionedB tBgB_nkl_, + STensorA tAsA_, STensorB tBsB_, + GTensorPartitionedSFA tAgSFA_mkl_, GTensorPartitionedSFB tBgSFB_nkl_, + STensorSFA tAsSFA_, STensorSFB tBsSFB_, + uint16_t mcast_mask_a_, uint16_t mcast_mask_b_, + uint16_t mcast_mask_sfa_, uint16_t mcast_mask_sfb_) + : k_tiles(k_tiles_) + , tAgA_mkl(tAgA_mkl_), tBgB_nkl(tBgB_nkl_) + , tAsA(tAsA_), tBsB(tBsB_) + , tAgSFA_mkl(tAgSFA_mkl_), tBgSFB_nkl(tBgSFB_nkl_) + , tAsSFA(tAsSFA_), tBsSFB(tBsSFB_) + , mcast_mask_a(mcast_mask_a_), mcast_mask_b(mcast_mask_b_) + , mcast_mask_sfa(mcast_mask_sfa_), mcast_mask_sfb(mcast_mask_sfb_) {} + }; + + template < + class TiledMma, + class FragmentA, class FragmentB, + class FragmentSFA, class FragmentSFB, + class SFATiledCopy, class SmemFrgSFA, class TmemFrgSFA, + class SFBTiledCopy, class SmemFrgSFB, class TmemFrgSFB + > + struct MmaParams { + TiledMma tiled_mma; + FragmentA tCrA; + FragmentB tCrB; + FragmentSFA tCtSFA; + FragmentSFB tCtSFB; + SFATiledCopy tiled_copy_s2t_SFA; + SmemFrgSFA thr_tCsSFA_s2t; + TmemFrgSFA thr_tCtSFA_s2t; + SFBTiledCopy tiled_copy_s2t_SFB; + SmemFrgSFB thr_tCsSFB_s2t; + TmemFrgSFB thr_tCtSFB_s2t; + + CUTLASS_DEVICE + MmaParams ( + TiledMma tiled_mma_, + FragmentA tCrA_, FragmentB tCrB_, FragmentSFA tCtSFA_, FragmentSFB tCtSFB_, + SFATiledCopy tiled_copy_s2t_SFA_, SmemFrgSFA thr_tCsSFA_s2t_, TmemFrgSFA thr_tCtSFA_s2t_, + SFBTiledCopy tiled_copy_s2t_SFB_, SmemFrgSFB thr_tCsSFB_s2t_, TmemFrgSFB thr_tCtSFB_s2t_) + : tiled_mma(tiled_mma_) + , tCrA(tCrA_), tCrB(tCrB_), tCtSFA(tCtSFA_), tCtSFB(tCtSFB_) + , tiled_copy_s2t_SFA(tiled_copy_s2t_SFA_), thr_tCsSFA_s2t(thr_tCsSFA_s2t_), thr_tCtSFA_s2t(thr_tCtSFA_s2t_) + , tiled_copy_s2t_SFB(tiled_copy_s2t_SFB_), thr_tCsSFB_s2t(thr_tCsSFB_s2t_), thr_tCtSFB_s2t(thr_tCtSFB_s2t_) {} + }; + + // Host side kernel arguments + struct Arguments { + ArrayElementA const* ptr_A{nullptr}; + StrideA dA{}; + ArrayElementB const* ptr_B{nullptr}; + StrideB dB{}; + ElementSF const* ptr_SFA{nullptr}; + LayoutSFA layout_SFA{}; + ElementSF const* ptr_SFB{nullptr}; + LayoutSFB layout_SFB{}; + RuntimeDataTypeA runtime_data_type_a{}; + RuntimeDataTypeB runtime_data_type_b{}; + }; + + // Device side kernel params + struct Params { + using ClusterLayout_VMNK = + decltype(tiled_divide(make_layout(conditional_return(make_shape(uint32_t(0), uint32_t(0), Int<1>{}), + ClusterShape{})), make_tile(typename TiledMma::AtomThrID{}))); + + using ClusterLayoutSfb_VMNK = + decltype(tiled_divide(make_layout(conditional_return(make_shape(uint32_t(0), uint32_t(0), Int<1>{}), + ClusterShape{})), make_tile(typename TiledMMA_SF::AtomThrID{}))); + + using TMA_A = decltype(make_tma_atom_A_sm100( + GmemTiledCopyA{}, + make_tensor(recast_ptr(nullptr), repeat_like(StrideA{}, int32_t(0)), StrideA{}), + SmemLayoutA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + ClusterLayout_VMNK{}) + ); + + using TMA_B = decltype(make_tma_atom_B_sm100( + GmemTiledCopyB{}, + make_tensor(recast_ptr(nullptr), repeat_like(StrideB{}, int32_t(0)), StrideB{}), + SmemLayoutB{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + ClusterLayout_VMNK{}) + ); + + using TMA_SFA = decltype(make_tma_atom_A_sm100( + GmemTiledCopySFA{}, + make_tensor(static_cast(nullptr), LayoutSFA{}), + SmemLayoutSFA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + ClusterLayout_VMNK{}) + ); + + using TMA_SFB = decltype(make_tma_atom_B_sm100( + GmemTiledCopySFB{}, + make_tensor(static_cast(nullptr), LayoutSFB{}), + SmemLayoutSFB{}(_,_,_,cute::Int<0>{}), + TileShape_SF{}, + TiledMMA_SF{}, + ClusterLayoutSfb_VMNK{}) + ); + + TMA_A tma_load_a; + TMA_B tma_load_b; + TMA_SFA tma_load_sfa; + TMA_SFB tma_load_sfb; + TMA_A tma_load_a_fallback; + TMA_B tma_load_b_fallback; + TMA_SFA tma_load_sfa_fallback; + TMA_SFB tma_load_sfb_fallback; + LayoutSFA layout_SFA; + LayoutSFB layout_SFB; + dim3 cluster_shape_fallback; + RuntimeDataTypeA runtime_data_type_a; + RuntimeDataTypeB runtime_data_type_b; + }; + + CUTLASS_DEVICE + CollectiveMma(Params const& params, ClusterShape cluster_shape, uint32_t block_rank_in_cluster) + : cluster_shape_(cluster_shape) + , block_rank_in_cluster_(block_rank_in_cluster) + , layout_SFA_(params.layout_SFA) + , layout_SFB_(params.layout_SFB) + , runtime_data_type_a_(params.runtime_data_type_a) + , runtime_data_type_b_(params.runtime_data_type_b) { + if constexpr (IsDynamicCluster) { + const bool is_fallback_cluster = (cute::size<0>(cluster_shape_) == params.cluster_shape_fallback.x && + cute::size<1>(cluster_shape_) == params.cluster_shape_fallback.y); + observed_tma_load_a_ = is_fallback_cluster ? ¶ms.tma_load_a_fallback : ¶ms.tma_load_a; + observed_tma_load_b_ = is_fallback_cluster ? ¶ms.tma_load_b_fallback : ¶ms.tma_load_b; + observed_tma_load_sfa_ = is_fallback_cluster ? ¶ms.tma_load_sfa_fallback : ¶ms.tma_load_sfa; + observed_tma_load_sfb_ = is_fallback_cluster ? ¶ms.tma_load_sfb_fallback : ¶ms.tma_load_sfb; + } + else { + observed_tma_load_a_ = ¶ms.tma_load_a; + observed_tma_load_b_ = ¶ms.tma_load_b; + observed_tma_load_sfa_ = ¶ms.tma_load_sfa; + observed_tma_load_sfb_ = ¶ms.tma_load_sfb; + } + } + + template + static constexpr Params + to_underlying_arguments( + ProblemShape const& problem_shape, + Arguments const& args, + [[maybe_unused]] void* workspace, + cutlass::KernelHardwareInfo const& hw_info = cutlass::KernelHardwareInfo{}) { + + // Optionally append 1s until problem shape is rank-4 (MNKL), in case it is only rank-3 (MNK) + auto problem_shape_MNKL = append<4>(problem_shape, 1); + auto [M,N,K,L] = problem_shape_MNKL; + + auto ptr_A = recast_ptr(args.ptr_A); + auto ptr_B = recast_ptr(args.ptr_B); + + Tensor tensor_a = make_tensor(ptr_A, make_layout(make_shape(M,K,L), args.dA)); + Tensor tensor_b = make_tensor(ptr_B, make_layout(make_shape(N,K,L), args.dB)); + auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape); + + // Cluster layout for TMA construction + auto cluster_layout_vmnk = tiled_divide(make_layout(cluster_shape), make_tile(typename TiledMma::AtomThrID{})); + auto cluster_shape_fallback = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape_fallback); + auto cluster_layout_vmnk_fallback = tiled_divide(make_layout(cluster_shape_fallback), make_tile(typename TiledMma::AtomThrID{})); + Tensor tensor_sfa = make_tensor(args.ptr_SFA, args.layout_SFA); + Tensor tensor_sfb = make_tensor(args.ptr_SFB, args.layout_SFB); + + // Cluster layout for TMA construction of SFB + auto cluster_layout_sfb_vmnk = tiled_divide(make_layout(cluster_shape), make_tile(typename TiledMMA_SF::AtomThrID{})); + auto cluster_layout_sfb_vmnk_fallback = tiled_divide(make_layout(cluster_shape_fallback), make_tile(typename TiledMMA_SF::AtomThrID{})); + + typename Params::TMA_A tma_load_a = make_tma_atom_A_sm100( + GmemTiledCopyA{}, + tensor_a, + SmemLayoutA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk); + + typename Params::TMA_B tma_load_b = make_tma_atom_B_sm100( + GmemTiledCopyB{}, + tensor_b, + SmemLayoutB{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk); + + typename Params::TMA_A tma_load_a_fallback = make_tma_atom_A_sm100( + GmemTiledCopyA{}, + tensor_a, + SmemLayoutA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk_fallback); + + typename Params::TMA_B tma_load_b_fallback = make_tma_atom_B_sm100( + GmemTiledCopyB{}, + tensor_b, + SmemLayoutB{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk_fallback); + + typename Params::TMA_SFA tma_load_sfa = make_tma_atom_A_sm100( + GmemTiledCopySFA{}, + tensor_sfa, + SmemLayoutSFA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk); + + typename Params::TMA_SFB tma_load_sfb = make_tma_atom_B_sm100( + GmemTiledCopySFB{}, + tensor_sfb, + SmemLayoutSFB{}(_,_,_,cute::Int<0>{}), + TileShape_SF{}, + TiledMMA_SF{}, + cluster_layout_sfb_vmnk); + + typename Params::TMA_SFA tma_load_sfa_fallback = make_tma_atom_A_sm100( + GmemTiledCopySFA{}, + tensor_sfa, + SmemLayoutSFA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk_fallback); + + typename Params::TMA_SFB tma_load_sfb_fallback = make_tma_atom_B_sm100( + GmemTiledCopySFB{}, + tensor_sfb, + SmemLayoutSFB{}(_,_,_,cute::Int<0>{}), + TileShape_SF{}, + TiledMMA_SF{}, + cluster_layout_sfb_vmnk_fallback); + + return { + tma_load_a, + tma_load_b, + tma_load_sfa, + tma_load_sfb, + tma_load_a_fallback, + tma_load_b_fallback, + tma_load_sfa_fallback, + tma_load_sfb_fallback, + args.layout_SFA, + args.layout_SFB, + hw_info.cluster_shape_fallback, + args.runtime_data_type_a, + args.runtime_data_type_b + }; + } + + template + static bool + can_implement( + ProblemShape const& problem_shape, + [[maybe_unused]] Arguments const& args) { + auto problem_shape_MNKL = append<4>(problem_shape, 1); + auto [M,N,K,L] = problem_shape_MNKL; + + constexpr int tma_alignment_bits_A = cutlass::detail::get_input_alignment_bits(); + constexpr int tma_alignment_bits_B = cutlass::detail::get_input_alignment_bits(); + + bool implementable = true; + constexpr int min_tma_aligned_elements_A = tma_alignment_bits_A / cute::sizeof_bits::value; + implementable = implementable && cutlass::detail::check_alignment(cute::make_shape(M,K,L), StrideA{}); + constexpr int min_tma_aligned_elements_B = tma_alignment_bits_B / cute::sizeof_bits::value; + implementable = implementable && cutlass::detail::check_alignment(cute::make_shape(N,K,L), StrideB{}); + + // Check for SFA SFB layout requirement + const auto layout_sfa_ref = take<0,2>(Sm1xxBlkScaledConfig::tile_atom_to_shape_SFA(problem_shape_MNKL)); + const auto layout_sfb_ref = take<0,2>(Sm1xxBlkScaledConfig::tile_atom_to_shape_SFB(problem_shape_MNKL)); + + implementable = implementable && (layout_sfa_ref == take<0,2>(args.layout_SFA)); + if (!implementable) { + CUTLASS_TRACE_HOST(" CAN IMPLEMENT: layout_SFA mismatch, layout_SFA needs to be K-major\n"); + } + + implementable = implementable && (layout_sfb_ref == take<0,2>(args.layout_SFB)); + if (!implementable) { + CUTLASS_TRACE_HOST(" CAN IMPLEMENT: layout_SFB mismatch, layout_SFB needs to be K-major\n"); + } + + if (!implementable) { + CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n"); + } + return implementable; + } + + /// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance + CUTLASS_DEVICE void + prefetch_tma_descriptors() { + cute::prefetch_tma_descriptor(observed_tma_load_a_->get_tma_descriptor()); + cute::prefetch_tma_descriptor(observed_tma_load_b_->get_tma_descriptor()); + cute::prefetch_tma_descriptor(observed_tma_load_sfa_->get_tma_descriptor()); + cute::prefetch_tma_descriptor(observed_tma_load_sfb_->get_tma_descriptor()); + } + + /// Construct A Single Stage's Accumulator Shape + CUTLASS_DEVICE static + auto + partition_accumulator_shape() { + auto acc_shape = partition_shape_C(TiledMma{}, take<0,2>(TileShape{})); // ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N) + + return acc_shape; + } + + template + CUTLASS_DEVICE static + auto + slice_accumulator(TmemStorage tmem_storage, int stage) { + return cute::make_tuple(tmem_storage.accumulators(_,_,_,stage)); + } + + template + CUTLASS_DEVICE static + auto + init_tmem_tensors(EpilogueTile epi_tile) { + TiledMma tiled_mma; + auto acc_shape = partition_accumulator_shape(); + // ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N,ACC_PIPE) where ACC_PIPE=2 so we can double buffer our accumulators for mainloop and epilogue. + Tensor accumulators = cutlass::detail::make_sm100_accumulator( + tiled_mma, acc_shape, EpilogueTile{}); + Tensor tCtSFA = make_tensor(shape(SmemLayoutAtomSFA{})); + Tensor tCtSFB = make_tensor(shape(SmemLayoutAtomSFB{})); + + TmemStorage tmem_storage; + tmem_storage.accumulators = accumulators; + tmem_storage.tCtSFA = tCtSFA; + tmem_storage.tCtSFB = tCtSFB; + + return tmem_storage; + } + + template + CUTLASS_DEVICE static + void + set_tmem_offsets(TmemStorage& tmem_storage, uint32_t tmem_base_addr) { + tmem_storage.accumulators.data() = tmem_base_addr; + tmem_storage.tCtSFA.data() = tmem_storage.accumulators.data().get() + cutlass::detail::find_tmem_tensor_col_offset(tmem_storage.accumulators); + tmem_storage.tCtSFB.data() = tmem_storage.tCtSFA.data().get() + cutlass::detail::find_tmem_tensor_col_offset(tmem_storage.tCtSFA); + } + + /// Set up the data needed by this collective for load. + /// Return tuple element contain + /// gA_mkl - The tiled tma tensor for input A + /// gB_nkl - The tiled tma tensor for input B + /// tAgA_mkl - partitioned gmem tensor for A + /// tBgB_nkl - partitioned gmem tensor for B + /// tAsA - partitioned smem tensor for A + /// tBsB - partitioned smem tensor for B + /// tAgSFA_mkl - partitioned gmem tensor for SFA + /// tBgSFB_nkl - partitioned gmem tensor for SFB + /// tAsSFA - partitioned tmem tensor for SFA + /// tAsSFB - partitioned tmem tensor for SFB + /// mcast_mask_a - tma multicast mask for A + /// mcast_mask_b - tma multicast mask for B + /// mcast_mask_sfa - tma multicast mask for SFA + /// mcast_mask_sfb - tma multicast mask for SFB + template + CUTLASS_DEVICE auto + load_init( + ProblemShape_MNKL const& problem_shape_MNKL, + TensorStorage& shared_tensors) const { + using X = Underscore; + + // Separate out problem shape for convenience + auto [M,N,K,L] = problem_shape_MNKL; + + // Represent the full tensors -- get these from TMA + Tensor mA_mkl = observed_tma_load_a_->get_tma_tensor(make_shape(M,K,L)); + Tensor mB_nkl = observed_tma_load_b_->get_tma_tensor(make_shape(N,K,L)); + + // Tile the tensors and defer the slice + Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M, BLK_K, m, k, l) + Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N, BLK_K, n, k, l) + + // Represent the full tensor of Scale factors + Tensor mSFA_mkl = observed_tma_load_sfa_->get_tma_tensor(shape(layout_SFA_)); + auto mSFB_nkl = [=](){ + if constexpr (IsCtaN192) { + Tensor mSFB_tmp = observed_tma_load_sfb_->get_tma_tensor(shape(layout_SFB_)); + auto x = stride<0,1>(mSFB_tmp); + auto y = ceil_div(shape<0,1>(mSFB_tmp), 4); + auto new_shape = make_shape (make_shape( shape<0,0>(mSFB_tmp), + make_shape( make_shape(_2{}, _2{}), y)), shape<1>(mSFB_tmp), shape<2>(mSFB_tmp)); + auto new_stride = make_stride(make_stride(stride<0,0>(mSFB_tmp), + make_stride(make_stride( x, x), x*3)), stride<1>(mSFB_tmp), stride<2>(mSFB_tmp)); + return make_tensor(mSFB_tmp.data(), make_layout(new_shape, new_stride)); + } + else if constexpr (IsCtaN64) { + Tensor mSFB_tmp = observed_tma_load_sfb_->get_tma_tensor(shape(layout_SFB_)); + auto new_shape = make_shape(make_shape(shape<0,0>(mSFB_tmp), + make_shape(_2{} , shape<0,1>(mSFB_tmp))), shape<1>(mSFB_tmp), shape<2>(mSFB_tmp)); + auto new_stride = make_stride(make_stride(stride<0,0>(mSFB_tmp), + make_stride(_0{}, stride<0,1>(mSFB_tmp))), stride<1>(mSFB_tmp), stride<2>(mSFB_tmp)); + return make_tensor(mSFB_tmp.data(), make_layout(new_shape, new_stride)); + } + else { + return observed_tma_load_sfb_->get_tma_tensor(shape(layout_SFB_)); + } + }(); + + Tensor gSFA_mkl = local_tile(mSFA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (TILE_M,TILE_K,m,k,l) + Tensor gSFB_nkl = local_tile(mSFB_nkl, TileShape_SF{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (TILE_N,TILE_K,n,k,l) + + // Partition for this CTA + ThrMMA cta_mma = TiledMma{}.get_slice(blockIdx.x % size(typename TiledMma::AtomThrID{})); + + Tensor tCgA_mkl = cta_mma.partition_A(gA_mkl); // (MMA, MMA_M, MMA_K, m, k, l) + Tensor tCgB_nkl = cta_mma.partition_B(gB_nkl); // (MMA, MMA_N, MMA_K, n, k, l) + + Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (MMA,MMA_M,MMA_K,PIPE) + Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (MMA,MMA_N,MMA_K,PIPE) + + ThrMMA cta_mma_sfb = TiledMMA_SF{}.get_slice(blockIdx.x % size(typename TiledMMA_SF::AtomThrID{})); + Tensor tCgSFA_mkl = cta_mma.partition_A(gSFA_mkl); // (MMA, MMA_M, MMA_K, m, k, l) + Tensor tCgSFB_nkl = cta_mma_sfb.partition_B(gSFB_nkl); // (MMA, MMA_N, MMA_K, n, k, l) + + Tensor sSFA = make_tensor(make_smem_ptr(shared_tensors.smem_SFA.begin()), SmemLayoutSFA{}); + Tensor sSFB = make_tensor(make_smem_ptr(shared_tensors.smem_SFB.begin()), SmemLayoutSFB{}); + + // Define the CTA-in-cluster Layout and Coord + Layout cta_layout_mnk = make_layout(cluster_shape_); + Layout cta_layout_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma::AtomThrID{})); + auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(block_rank_in_cluster_); + + Layout cta_layout_sfb_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMMA_SF::AtomThrID{})); + auto cta_coord_sfb_vmnk = cta_layout_sfb_vmnk.get_flat_coord(block_rank_in_cluster_); + + // Project the cta_layout for tma_a along the n-modes + auto [tAgA_mkl, tAsA] = tma_partition(*observed_tma_load_a_, + get<2>(cta_coord_vmnk), make_layout(size<2>(cta_layout_vmnk)), + group_modes<0,3>(sA), group_modes<0,3>(tCgA_mkl)); + + // Project the cta_layout for tma_b along the m-modes + auto [tBgB_nkl, tBsB] = tma_partition(*observed_tma_load_b_, + get<1>(cta_coord_vmnk), make_layout(size<1>(cta_layout_vmnk)), + group_modes<0,3>(sB), group_modes<0,3>(tCgB_nkl)); + + // Project the cta_layout for tma_a along the n-modes + auto [tAgSFA_mkl, tAsSFA] = tma_partition(*observed_tma_load_sfa_, + get<2>(cta_coord_vmnk), make_layout(size<2>(cta_layout_vmnk)), + group_modes<0,3>(sSFA), group_modes<0,3>(tCgSFA_mkl)); + + // Project the cta_layout for tma_b along the m-modes + auto [tBgSFB_nkl, tBsSFB] = tma_partition(*observed_tma_load_sfb_, + get<1>(cta_coord_sfb_vmnk), make_layout(size<1>(cta_layout_sfb_vmnk)), + group_modes<0,3>(sSFB), group_modes<0,3>(tCgSFB_nkl)); + + // TMA Multicast Masks + uint16_t mcast_mask_a = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk); + uint16_t mcast_mask_b = create_tma_multicast_mask<1>(cta_layout_vmnk, cta_coord_vmnk); + uint16_t mcast_mask_sfa = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk); + uint16_t mcast_mask_sfb = create_tma_multicast_mask<1>(cta_layout_sfb_vmnk, cta_coord_sfb_vmnk); + + return LoadParams{ + size<3>(gA_mkl), // for scheduler + tAgA_mkl, tBgB_nkl, tAsA, tBsB, // for input tensor values + tAgSFA_mkl, tBgSFB_nkl, tAsSFA, tBsSFB, // for input scale factor tensor values + mcast_mask_a, mcast_mask_b, mcast_mask_sfa, mcast_mask_sfb}; // multicast masks + } + + /// Set up the data needed by this collective for mma compute. + template + CUTLASS_DEVICE auto + mma_init( + TmemStorage tmem_storage, + TensorStorage& shared_tensors) const { + + // Allocate "fragments/descriptors" for A and B matrices + Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (BLK_M,BLK_K,PIPE) + Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (BLK_N,BLK_K,PIPE) + + // Allocate "fragments/descriptors" for A and B matrices + Tensor tCrA = TiledMma::make_fragment_A(sA); // (MMA,MMA_M,MMA_K,PIPE) + Tensor tCrB = TiledMma::make_fragment_B(sB); // (MMA,MMA_N,MMA_K,PIPE) + + CUTE_STATIC_ASSERT_V(Int{} == size<3>(sA)); // PIPE + CUTE_STATIC_ASSERT_V(Int{} == size<3>(sB)); // PIPE + + // + // Scale Factor + // + Tensor tCtSFA = tmem_storage.tCtSFA; + Tensor tCtSFB = tmem_storage.tCtSFB; + // Setup smem descriptors for UTCCP + Tensor tCsSFA = make_tensor(make_smem_ptr(shared_tensors.smem_SFA.begin()), SmemLayoutSFA{}); + Tensor tCsSFB = make_tensor(make_smem_ptr(shared_tensors.smem_SFB.begin()), SmemLayoutSFB{}); + + // Make SMEM and TMEM tensors compact removing the zero strides to eliminate unnecessary copy instructions. + auto tCsSFA_compact = make_tensor(tCsSFA.data(), filter_zeros(tCsSFA.layout())); + auto tCtSFA_compact = make_tensor(tCtSFA.data(), filter_zeros(tCtSFA.layout())); + auto tCsSFB_compact = make_tensor(tCsSFB.data(), filter_zeros(tCsSFB.layout())); + auto tCtSFB_compact = make_tensor(tCtSFB.data(), filter_zeros(tCtSFB.layout())); + + // Create the SMEM to TMEM copy operations based on the MMA atom used (1CTA vs 2CTA) + using AtomThrID = typename TiledMma::AtomThrID; + using UtccpOp = cute::conditional_t<(decltype(cute::size(AtomThrID{}) == Int<2>{})::value), + SM100_UTCCP_4x32dp128bit_2cta, SM100_UTCCP_4x32dp128bit_1cta>; + + // Both SM100_UTCCP_4x32dp128bit_{1,2}cta replicate one SMEM core matrix across + // 4 TMEM core matrices (the "4x" in the name). + constexpr int UtccpBroadcastFactor = 4; + auto tiled_copy_s2t_SFA = make_utccp_copy(UtccpOp{}, tCtSFA_compact); + auto tiled_copy_s2t_SFB = make_utccp_copy(UtccpOp{}, tCtSFB_compact); + + // make_utccp_copy() builds its tiler from TMEM alone, so the SMEM source must have + // the replication it implies spelled out before partition_S() -- see + // sm107_append_s2t_broadcast_mode()'s comment above. + auto tCsSFA_compact_bcast = sm107_append_s2t_broadcast_mode(tCsSFA_compact); + auto tCsSFB_compact_bcast = sm107_append_s2t_broadcast_mode(tCsSFB_compact); + + auto thr_copy_s2t_SFA = tiled_copy_s2t_SFA.get_slice(0); + auto thr_tCsSFA_compact_s2t_ = thr_copy_s2t_SFA.partition_S(tCsSFA_compact_bcast); + // SMEM to TMEM copy operation requires source SMEM operand to be an SMEM descriptor + auto thr_tCsSFA_compact_s2t = get_utccp_smem_desc_tensor(thr_tCsSFA_compact_s2t_); + auto thr_tCtSFA_compact_s2t = thr_copy_s2t_SFA.partition_D(tCtSFA_compact); + + auto thr_copy_s2t_SFB = tiled_copy_s2t_SFB.get_slice(0); + auto thr_tCsSFB_compact_s2t_ = thr_copy_s2t_SFB.partition_S(tCsSFB_compact_bcast); + // SMEM to TMEM copy operation requires source SMEM operand to be an SMEM descriptor + auto thr_tCsSFB_compact_s2t = get_utccp_smem_desc_tensor(thr_tCsSFB_compact_s2t_); + auto thr_tCtSFB_compact_s2t = thr_copy_s2t_SFB.partition_D(tCtSFB_compact); + + TiledMma tiled_mma; + + if constexpr (IsRuntimeDataType) { + // Update instruction descriptor according to runtime argument. + // Applying bitmask (0b111) to help compiler deduce that the conversion and assignment are safe. + tiled_mma.idesc_.a_format_ = uint8_t(runtime_data_type_a_) & 0b111; + tiled_mma.idesc_.b_format_ = uint8_t(runtime_data_type_b_) & 0b111; + } + + return MmaParams{ + tiled_mma, + tCrA, tCrB, tCtSFA, tCtSFB, + tiled_copy_s2t_SFA, thr_tCsSFA_compact_s2t, thr_tCtSFA_compact_s2t, + tiled_copy_s2t_SFB, thr_tCsSFB_compact_s2t, thr_tCtSFB_compact_s2t}; + } + + /// Perform a collective-scoped matrix multiply-accumulate + /// Producer Perspective + template < + class LoadParams, + class TileCoordMNKL, + class KTileIterator + > + CUTLASS_DEVICE auto + load( + MainloopPipeline mainloop_pipeline, + MainloopPipelineState mainloop_pipe_producer_state, + LoadParams const& load_inputs, + TileCoordMNKL const& cta_coord_mnkl, + KTileIterator k_tile_iter, int k_tile_count) { + + auto [unused_k_tiles, + tAgA_mkl, tBgB_nkl, tAsA, tBsB, + tAgSFA_mkl, tBgSFB_nkl, tAsSFA, tBsSFB, + mcast_mask_a, mcast_mask_b, mcast_mask_sfa, mcast_mask_sfb] = load_inputs; + + // slice out the work coord from partitioned tensors + Tensor tAgA = tAgA_mkl(_, get<0>(cta_coord_mnkl) / size(typename TiledMma::AtomThrID{}), _, get<3>(cta_coord_mnkl)); + Tensor tBgB = tBgB_nkl(_, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl)); + Tensor tAgSFA = tAgSFA_mkl(_, get<0>(cta_coord_mnkl) / size(typename TiledMma::AtomThrID{}), _, get<3>(cta_coord_mnkl)); + Tensor tBgSFB = tBgSFB_nkl(_, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl)); + + auto barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state); + + // Issue the Mainloop loads + CUTLASS_PRAGMA_NO_UNROLL + while (k_tile_count > 0) { + // LOCK mainloop_pipe_producer_state for _writing_ + mainloop_pipeline.producer_acquire(mainloop_pipe_producer_state, barrier_token); + // Note: We don't synchronize the sf_pipeline for "Buffer_Empty". We use mainloop pipeline + // to do the synchronization at once. + + using BarrierType = typename MainloopPipeline::ProducerBarrierType; + BarrierType* tma_barrier = mainloop_pipeline.producer_get_barrier(mainloop_pipe_producer_state); + + int write_stage = mainloop_pipe_producer_state.index(); + ++mainloop_pipe_producer_state; + barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state); + + if (cute::elect_one_sync()) { + copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage)); + copy(observed_tma_load_b_->with(*tma_barrier, mcast_mask_b), tBgB(_,*k_tile_iter), tBsB(_,write_stage)); + copy(observed_tma_load_sfa_->with(*tma_barrier, mcast_mask_sfa), tAgSFA(_,*k_tile_iter), tAsSFA(_,write_stage)); + copy(observed_tma_load_sfb_->with(*tma_barrier, mcast_mask_sfb), tBgSFB(_,*k_tile_iter), tBsSFB(_,write_stage)); + } + + --k_tile_count; + ++k_tile_iter; + } + + return cute::make_tuple(mainloop_pipe_producer_state, k_tile_iter); + } + + /// Perform a Producer Epilogue to prevent early exit of ctas in a Cluster + CUTLASS_DEVICE void + load_tail(MainloopPipeline mainloop_pipeline, MainloopPipelineState mainloop_pipe_producer_state) { + // Issue the epilogue waits + // This helps avoid early exit of ctas in Cluster + // Waits for all stages to either be released (all + // Consumer UNLOCKs), or if the stage was never used + // then would just be acquired since the phase was + // still inverted from make_producer_start_state + mainloop_pipeline.producer_tail(mainloop_pipe_producer_state); + } + + /// Perform a collective-scoped matrix multiply-accumulate + /// Consumer Perspective + template < + class AccumulatorPipeline, + class FrgEngine, class FrgLayout, + class MmaParams, + class CtaTileCoord + > + CUTLASS_DEVICE auto + mma(cute::tuple pipelines, + cute::tuple pipeline_states, + cute::tuple> const& accumulators_pair, + MmaParams const& mma_inputs, + CtaTileCoord cta_tile_coord, + int k_tile_count + ) { + static_assert(is_tmem::value, "Accumulator must be tmem resident."); + static_assert(rank(FrgLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA, MMA_M, MMA_N)"); + + auto accumulators = get<0>(accumulators_pair); + auto [tiled_mma, + tCrA, tCrB, tCtSFA, tCtSFB, + tiled_copy_s2t_SFA, thr_tCsSFA_s2t, + thr_tCtSFA_s2t, tiled_copy_s2t_SFB, + thr_tCsSFB_s2t, thr_tCtSFB_s2t] = mma_inputs; + + auto [mainloop_pipeline, accumulator_pipeline] = pipelines; + auto [mainloop_pipe_consumer_state, accumulator_pipe_producer_state] = pipeline_states; + + auto tCtSFB_mma = [tCtSFB = tCtSFB, cta_tile_coord]() { + if constexpr (IsCtaN192) { + // If this is an ODD tile, shift the TMEM start address for N=192 case by two words (ignores first 64 columns of SFB) + auto tCtSFB_tmp = tCtSFB; + if (size<1>(cta_tile_coord) % 2 == 1) { + tCtSFB_tmp.data() = tCtSFB_tmp.data().get() + 2; + } + return tCtSFB_tmp; + } + else if constexpr (IsCtaN64) { + // Move in increments of 64 columns of SFB + auto tCtSFB_tmp = tCtSFB; + tCtSFB_tmp.data() = tCtSFB_tmp.data().get() + (size<1>(cta_tile_coord) % 2) * 2; + return tCtSFB_tmp; + } + else { + return tCtSFB; + } + }(); + + uint32_t skip_wait = k_tile_count <= 0; + auto barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait); + + // + // PIPELINED MAIN LOOP + // + tiled_mma.accumulate_ = UMMA::ScaleOut::Zero; + // Wait for tmem accumulator buffer to become empty with a flipped phase + accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state); + + CUTLASS_PRAGMA_NO_UNROLL + while (k_tile_count > 0) { + // WAIT on mainloop_pipe_consumer_state until its data are available + // (phase bit flips from mainloop_pipe_consumer_state.phase() value) + mainloop_pipeline.consumer_wait(mainloop_pipe_consumer_state, barrier_token); + + // Compute on k_tile + int read_stage = mainloop_pipe_consumer_state.index(); + // Save current mainlop pipeline read state + auto curr_mainloop_pipe_consumer_state = mainloop_pipe_consumer_state; + + // Advance mainloop_pipe + ++mainloop_pipe_consumer_state; + --k_tile_count; + skip_wait = k_tile_count <= 0; + // Peek at next iteration + barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait); + + if (cute::elect_one_sync()) { + copy(tiled_copy_s2t_SFA, thr_tCsSFA_s2t(_,_,_,_,read_stage), thr_tCtSFA_s2t); + copy(tiled_copy_s2t_SFB, thr_tCsSFB_s2t(_,_,_,_,read_stage), thr_tCtSFB_s2t); + } + + // Unroll the K mode manually so we can set scale C to 1 + CUTLASS_PRAGMA_UNROLL + for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) { + // (V,M) x (V,N) => (V,M,N) + if constexpr (WithBreuse) { + static_assert(size<1>(tCrA) == 2, + "The b-reuse feature expects size<1>(tCrA) == 2."); + static_assert(size<1>(tCrB) == 1, + "The b-reuse feature expects size<1>(tCrB) == 1."); + static_assert(size<1>(accumulators) == 2, + "The b-reuse feature expects size<1>(accumulators) == 2."); + + cute::gemm(tiled_mma.with(C{}), + make_zip_tensor(tCrA(_,0,k_block,read_stage), tCtSFA(_,0,k_block)), + make_zip_tensor(tCrB(_,0,k_block,read_stage), tCtSFB_mma(_,0,k_block)), + accumulators(_,0,0)); + cute::gemm(tiled_mma.with(C{}), + make_zip_tensor(tCrA(_,1,k_block,read_stage), tCtSFA(_,1,k_block)), + make_zip_tensor(tCrB(_,0,k_block,read_stage), tCtSFB_mma(_,0,k_block)), + accumulators(_,1,0)); + } + else { + cute::gemm(tiled_mma, + make_zip_tensor(tCrA(_,_,k_block,read_stage), tCtSFA(_,_,k_block)), + make_zip_tensor(tCrB(_,_,k_block,read_stage), tCtSFB_mma(_,_,k_block)), + accumulators); + } + tiled_mma.accumulate_ = UMMA::ScaleOut::One; + } + + mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state); + } + + return mainloop_pipe_consumer_state; + } + +protected: + + typename Params::TMA_A const* observed_tma_load_a_{nullptr}; + typename Params::TMA_B const* observed_tma_load_b_{nullptr}; + typename Params::TMA_SFA const* observed_tma_load_sfa_{nullptr}; + typename Params::TMA_SFB const* observed_tma_load_sfb_{nullptr}; + + LayoutSFA layout_SFA_; + LayoutSFB layout_SFB_; + RuntimeDataTypeA runtime_data_type_a_{}; + RuntimeDataTypeB runtime_data_type_b_{}; + + ClusterShape cluster_shape_; + uint32_t block_rank_in_cluster_; +}; + +///////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace cutlass::gemm::collective + +///////////////////////////////////////////////////////////////////////////////////////////////// diff --git a/include/cutlass/gemm/collective/sm107_mma_warpspecialized.hpp b/include/cutlass/gemm/collective/sm107_mma_warpspecialized.hpp new file mode 100644 index 0000000000..2c21aee247 --- /dev/null +++ b/include/cutlass/gemm/collective/sm107_mma_warpspecialized.hpp @@ -0,0 +1,700 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#pragma once + +#include "cutlass/cutlass.h" +#include "cutlass/detail/collective.hpp" +#include "cutlass/detail/cluster.hpp" +#include "cutlass/gemm/dispatch_policy.hpp" +#include "cutlass/numeric_types.h" +#include "cutlass/pipeline/pipeline.hpp" +#include "cutlass/gemm/gemm.h" +#include "cutlass/trace.h" +#include "cutlass/kernel_hardware_info.hpp" +#include "cutlass/detail/sm100_tmem_helper.hpp" + +#include "cute/algorithm/functional.hpp" +#include "cute/arch/cluster_sm90.hpp" +#include "cute/atom/mma_atom.hpp" +#include "cute/algorithm/gemm.hpp" +#include "cute/numeric/arithmetic_tuple.hpp" + +///////////////////////////////////////////////////////////////////////////////////////////////// + +namespace cutlass::gemm::collective { +using namespace cute; + +///////////////////////////////////////////////////////////////////////////////////////////////// + +// WarpSpecialized Mainloop for SM107 +// Both DMA Load and MMA methods of this class must be run by a single thread that's picked by elect_one. +template < + int Stages, + int SchedulerPipelineStageCount, + int AccumulatorPipelineStageCount, + class ClusterShape, // Static cluster shape or dynamic (int, int, _1) + bool WithBreuse_, + class TileShape_, // (MmaAtomShapeM, MmaAtomShapeN, TileK) + class ElementA_, + class StrideA_, + class ElementB_, + class StrideB_, + class TiledMma_, + class GmemTiledCopyA_, + class SmemLayoutAtomA_, + class SmemCopyAtomA_, + class TransformA_, + class GmemTiledCopyB_, + class SmemLayoutAtomB_, + class SmemCopyAtomB_, + class TransformB_ +> +struct CollectiveMma< + MainloopSm107TmaUmmaWarpSpecialized< + Stages, + SchedulerPipelineStageCount, + AccumulatorPipelineStageCount, + ClusterShape, + WithBreuse_, + arch::Sm107 + >, + TileShape_, + ElementA_, StrideA_, + ElementB_, StrideB_, + TiledMma_, + GmemTiledCopyA_, SmemLayoutAtomA_, SmemCopyAtomA_, TransformA_, + GmemTiledCopyB_, SmemLayoutAtomB_, SmemCopyAtomB_, TransformB_> +{ + // + // Type Aliases + // + using TiledMma = TiledMma_; + using AtomThrShapeMNK = Shape(typename TiledMma::ThrLayoutVMNK{})), _1, _1>; + + static constexpr bool WithBreuse = WithBreuse_; + using DispatchPolicy = MainloopSm107TmaUmmaWarpSpecialized< + Stages, + SchedulerPipelineStageCount, + AccumulatorPipelineStageCount, + ClusterShape, + WithBreuse, + arch::Sm107 + >; + using TileShape = TileShape_; + + static constexpr bool IsDynamicCluster = not cute::is_static_v; + + CUTE_STATIC_ASSERT_V(evenly_divides(TileShape{}, tile_shape(TiledMma{})), + "Static cluster shape used: TileShape should be evenly divided by TiledMma"); + + using CtaShape_MNK = decltype(shape_div(TileShape{}, AtomThrShapeMNK{})); + + // Define A and B block shapes for reduced size TMA_LOADs + using MmaShapeA_MK = decltype(partition_shape_A(TiledMma{}, make_shape(size<0>(TileShape{}), size<2>(TileShape{})))); + using MmaShapeB_NK = decltype(partition_shape_B(TiledMma{}, make_shape(size<1>(TileShape{}), size<2>(TileShape{})))); + + using ElementA = ElementA_; + using ElementAMma = typename TiledMma::ValTypeA; + using StrideA = StrideA_; + using ElementB = ElementB_; + using ElementBMma = typename TiledMma::ValTypeB; + using StrideB = StrideB_; + + static constexpr bool IsRuntimeDataTypeA = cutlass::gemm::collective::detail::is_sm10x_runtime_f8f6f4(); + static constexpr bool IsRuntimeDataTypeB = cutlass::gemm::collective::detail::is_sm10x_runtime_f8f6f4(); + + static_assert((IsRuntimeDataTypeA && IsRuntimeDataTypeB) || + (!IsRuntimeDataTypeA && !IsRuntimeDataTypeB), + "ElementA and ElementB should be both runtime or both static."); + + static constexpr bool IsRuntimeDataType = IsRuntimeDataTypeA && IsRuntimeDataTypeB; + + using ElementAccumulator = typename TiledMma::ValTypeC; + using GmemTiledCopyA = GmemTiledCopyA_; + using GmemTiledCopyB = GmemTiledCopyB_; + using SmemLayoutAtomA = SmemLayoutAtomA_; + using SmemLayoutAtomB = SmemLayoutAtomB_; + using SmemCopyAtomA = SmemCopyAtomA_; + using SmemCopyAtomB = SmemCopyAtomB_; + using TransformA = TransformA_; + using TransformB = TransformB_; + using ArchTag = typename DispatchPolicy::ArchTag; + + using MainloopPipeline = cutlass::PipelineTmaUmmaAsync< + DispatchPolicy::Stages, + ClusterShape, + AtomThrShapeMNK>; + using MainloopPipelineState = typename MainloopPipeline::PipelineState; + + static_assert(rank(SmemLayoutAtomA{}) == 2, "SmemLayoutAtomA must be rank 2 (M,K)"); + static_assert(((size<0,0>(MmaShapeA_MK{}) * size<1>(MmaShapeA_MK{})) % size<0>(SmemLayoutAtomA{})) == 0, + "SmemLayoutAtom must evenly divide tile shape."); + static_assert(((size<0,1>(MmaShapeA_MK{}) * size<2>(MmaShapeA_MK{})) % size<1>(SmemLayoutAtomA{})) == 0, + "SmemLayoutAtom must evenly divide tile shape."); + static_assert(cute::is_void_v, + "SM107 UMMA cannot have a non-void copy atom for smem sourced instructions."); + + static_assert(rank(SmemLayoutAtomB{}) == 2, "SmemLayoutAtomB must be rank 2 (N,K)"); + static_assert(((size<0,0>(MmaShapeB_NK{}) * size<1>(MmaShapeB_NK{})) % size<0>(SmemLayoutAtomB{})) == 0, + "SmemLayoutAtom must evenly divide tile shape."); + static_assert(((size<0,1>(MmaShapeB_NK{}) * size<2>(MmaShapeB_NK{})) % size<1>(SmemLayoutAtomB{})) == 0, + "SmemLayoutAtom must evenly divide tile shape."); + static_assert(cute::is_void_v, + "SM107 UMMA cannot have a non-void copy atom for smem sourced instructions."); + + static_assert( + (size(AtomThrShapeMNK{}) == 1 && + (cute::is_same_v || cute::is_same_v)) || + (size(AtomThrShapeMNK{}) == 2 && + (cute::is_same_v || cute::is_same_v)), + "GmemTiledCopy - invalid TMA copy atom specified."); + static_assert( + (size(AtomThrShapeMNK{}) == 1 && + (cute::is_same_v || cute::is_same_v)) || + (size(AtomThrShapeMNK{}) == 2 && + (cute::is_same_v || cute::is_same_v)), + "GmemTiledCopy - invalid TMA copy atom specified."); + + using TmaInternalElementA = cute::conditional_t, cutlass::tfloat32_t, ElementAMma>; + using TmaInternalElementB = cute::conditional_t, cutlass::tfloat32_t, ElementBMma>; + + using SmemAllocTypeA = cute::conditional_t < 8, uint8_t, ElementAMma>; + using SmemAllocTypeB = cute::conditional_t < 8, uint8_t, ElementBMma>; + + using BitTypeElementA = cute::uint_bit_t>; + using BitTypeElementB = cute::uint_bit_t>; + + using ArrayElementA = cute::conditional_t; + using ArrayElementB = cute::conditional_t; + + using RuntimeDataTypeA = cute::conditional_t; + using RuntimeDataTypeB = cute::conditional_t; + + // Tile along K mode first before tiling over MN. PIPE mode last as usual. + // (MMA_TILE_M,MMA_TILE_K),MMA_M,MMA_K,PIPE) + using SmemLayoutA = decltype(UMMA::tile_to_mma_shape( + SmemLayoutAtomA{}, + append(MmaShapeA_MK{}, Int{}), + cute::conditional_t(), Step<_2,_1,_3>, Step<_1,_2,_3>>{})); + // (MMA_TILE_N,MMA_TILE_K),MMA_N,MMA_K,PIPE) + using SmemLayoutB = decltype(UMMA::tile_to_mma_shape( + SmemLayoutAtomB{}, + append(MmaShapeB_NK{}, Int{}), + cute::conditional_t(), Step<_2,_1,_3>, Step<_1,_2,_3>>{})); + + static_assert(DispatchPolicy::Stages >= 2, "Specialization requires Stages set to value 2 or more."); + static_assert(cute::is_base_of::value && + cute::is_base_of::value, + "MMA atom must source both A and B operand from smem_desc for this mainloop."); + + struct SharedStorage { + struct TensorStorage : cute::aligned_struct<128, _0> { + cute::ArrayEngine> smem_A; + cute::ArrayEngine> smem_B; + } tensors; + + using PipelineStorage = typename MainloopPipeline::SharedStorage; + PipelineStorage pipeline; + }; + + // Expose shared storage for tensors/pipelines separately to allow kernel layer to reorder them. + using TensorStorage = typename SharedStorage::TensorStorage; + using PipelineStorage = typename SharedStorage::PipelineStorage; + + // Only one thread issues the TMA and updates the barriers in a 2SM MMA, adjust bytes accordingly + static constexpr uint32_t TmaTransactionBytesA = + cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutA{})) * cute::sizeof_bits_v); + static constexpr uint32_t TmaTransactionBytesB = + cutlass::bits_to_bytes(size(AtomThrShapeMNK{}) * cosize(take<0,3>(SmemLayoutB{})) * cute::sizeof_bits_v); + static constexpr uint32_t TmaTransactionBytes = TmaTransactionBytesA + TmaTransactionBytesB; + + template + struct TmemStorage { + AccTensor accumulators; + }; + + template < + class KTileCount, + class GTensorPartitionedA, class GTensorPartitionedB, + class STensorA, class STensorB + > + struct LoadParams { + KTileCount k_tiles; + GTensorPartitionedA tAgA_mkl; + GTensorPartitionedB tBgB_nkl; + STensorA tAsA; + STensorB tBsB; + uint16_t mcast_mask_a; + uint16_t mcast_mask_b; + + CUTLASS_DEVICE + LoadParams ( + KTileCount k_tiles_, + GTensorPartitionedA tAgA_mkl_, GTensorPartitionedB tBgB_nkl_, + STensorA tAsA_, STensorB tBsB_, + uint16_t mcast_mask_a_, uint16_t mcast_mask_b_) + : k_tiles(k_tiles_) + , tAgA_mkl(tAgA_mkl_), tBgB_nkl(tBgB_nkl_) + , tAsA(tAsA_), tBsB(tBsB_) + , mcast_mask_a(mcast_mask_a_), mcast_mask_b(mcast_mask_b_) {} + }; + + template < + class TiledMma, + class FragmentA, class FragmentB + > + struct MmaParams { + TiledMma tiled_mma; + FragmentA tCrA; + FragmentB tCrB; + + CUTLASS_DEVICE + MmaParams ( + TiledMma tiled_mma_, + FragmentA tCrA_, FragmentB tCrB_) + : tiled_mma(tiled_mma_) + , tCrA(tCrA_), tCrB(tCrB_) {} + }; + + // Host side kernel arguments + struct Arguments { + ArrayElementA const* ptr_A{nullptr}; + StrideA dA{}; + ArrayElementB const* ptr_B{nullptr}; + StrideB dB{}; + RuntimeDataTypeA runtime_data_type_a{}; + RuntimeDataTypeB runtime_data_type_b{}; + }; + + // Device side kernel params + struct Params { + using ClusterLayout_VMNK = decltype(tiled_divide(make_layout(conditional_return(make_shape(uint32_t(0), uint32_t(0), Int<1>{}), ClusterShape{})), + make_tile(typename TiledMma::AtomThrID{}))); + + using TMA_A = decltype(make_tma_atom_A_sm100( + GmemTiledCopyA{}, + make_tensor(recast_ptr(nullptr), repeat_like(StrideA{}, int32_t(0)), StrideA{}), + SmemLayoutA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + ClusterLayout_VMNK{}) + ); + + using TMA_B = decltype(make_tma_atom_B_sm100( + GmemTiledCopyB{}, + make_tensor(recast_ptr(nullptr), repeat_like(StrideB{}, int32_t(0)), StrideB{}), + SmemLayoutB{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + ClusterLayout_VMNK{}) + ); + + TMA_A tma_load_a; + TMA_B tma_load_b; + TMA_A tma_load_a_fallback; + TMA_B tma_load_b_fallback; + dim3 cluster_shape_fallback; + RuntimeDataTypeA runtime_data_type_a; + RuntimeDataTypeB runtime_data_type_b; + }; + + CUTLASS_DEVICE + CollectiveMma(Params const& params, ClusterShape cluster_shape, uint32_t block_rank_in_cluster) + : cluster_shape_(cluster_shape) + , block_rank_in_cluster_(block_rank_in_cluster) + , runtime_data_type_a_(params.runtime_data_type_a) + , runtime_data_type_b_(params.runtime_data_type_b) { + if constexpr (IsDynamicCluster) { + const bool is_fallback_cluster = (cute::size<0>(cluster_shape_) == params.cluster_shape_fallback.x && + cute::size<1>(cluster_shape_) == params.cluster_shape_fallback.y); + observed_tma_load_a_ = is_fallback_cluster ? ¶ms.tma_load_a_fallback : ¶ms.tma_load_a; + observed_tma_load_b_ = is_fallback_cluster ? ¶ms.tma_load_b_fallback : ¶ms.tma_load_b; + } + else { + observed_tma_load_a_ = ¶ms.tma_load_a; + observed_tma_load_b_ = ¶ms.tma_load_b; + } + } + + template + static constexpr Params + to_underlying_arguments( + ProblemShape const& problem_shape, + Arguments const& args, + [[maybe_unused]] void* workspace, + cutlass::KernelHardwareInfo const& hw_info = cutlass::KernelHardwareInfo{}) { + + auto problem_shape_MNKL = append<4>(problem_shape, 1); + auto [M,N,K,L] = problem_shape_MNKL; + + auto ptr_A = recast_ptr(args.ptr_A); + auto ptr_B = recast_ptr(args.ptr_B); + + Tensor tensor_a = make_tensor(ptr_A, make_layout(make_shape(M,K,L), args.dA)); + Tensor tensor_b = make_tensor(ptr_B, make_layout(make_shape(N,K,L), args.dB)); + + auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape); + auto cluster_layout_vmnk = tiled_divide(make_layout(cluster_shape), make_tile(typename TiledMma::AtomThrID{})); + auto cluster_shape_fallback = cutlass::detail::select_cluster_shape(ClusterShape{}, hw_info.cluster_shape_fallback); + auto cluster_layout_vmnk_fallback = tiled_divide(make_layout(cluster_shape_fallback), make_tile(typename TiledMma::AtomThrID{})); + + typename Params::TMA_A tma_load_a = make_tma_atom_A_sm100( + GmemTiledCopyA{}, + tensor_a, + SmemLayoutA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk); + + typename Params::TMA_B tma_load_b = make_tma_atom_B_sm100( + GmemTiledCopyB{}, + tensor_b, + SmemLayoutB{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk); + + typename Params::TMA_A tma_load_a_fallback = make_tma_atom_A_sm100( + GmemTiledCopyA{}, + tensor_a, + SmemLayoutA{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk_fallback); + + typename Params::TMA_B tma_load_b_fallback = make_tma_atom_B_sm100( + GmemTiledCopyB{}, + tensor_b, + SmemLayoutB{}(_,_,_,cute::Int<0>{}), + TileShape{}, + TiledMma{}, + cluster_layout_vmnk_fallback); + + return { + tma_load_a, + tma_load_b, + tma_load_a_fallback, + tma_load_b_fallback, + hw_info.cluster_shape_fallback, + args.runtime_data_type_a, + args.runtime_data_type_b + }; + } + + template + static bool + can_implement( + ProblemShape const& problem_shape, + [[maybe_unused]] Arguments const& args) { + + auto problem_shape_MNKL = append<4>(problem_shape, 1); + auto [M,N,K,L] = problem_shape_MNKL; + + static constexpr bool IsF8F6F4 = detail::is_sm10x_f8f6f4_inputs(); + constexpr int tma_alignment_bits_A = cutlass::detail::get_input_alignment_bits(); + constexpr int tma_alignment_bits_B = cutlass::detail::get_input_alignment_bits(); + constexpr int min_tma_aligned_elements_A = tma_alignment_bits_A / cute::sizeof_bits::value; + constexpr int min_tma_aligned_elements_B = tma_alignment_bits_B / cute::sizeof_bits::value; + + bool implementable = true; + implementable = implementable && cutlass::detail::check_alignment(cute::make_shape(M,K,L), StrideA{}); + implementable = implementable && cutlass::detail::check_alignment(cute::make_shape(N,K,L), StrideB{}); + + if (!implementable) { + CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Problem Size doesn't meet the minimum alignment requirements for TMA.\n"); + } + return implementable; + } + + /// Issue Tma Descriptor Prefetch -- ideally from a single thread for best performance + CUTLASS_DEVICE void + prefetch_tma_descriptors() { + cute::prefetch_tma_descriptor(observed_tma_load_a_->get_tma_descriptor()); + cute::prefetch_tma_descriptor(observed_tma_load_b_->get_tma_descriptor()); + } + + /// Construct A Single Stage's Accumulator Shape + CUTLASS_DEVICE static + auto + partition_accumulator_shape() { + auto acc_shape = partition_shape_C(TiledMma{}, take<0,2>(TileShape{})); // ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N) + return acc_shape; + } + + template + CUTLASS_DEVICE static + auto + slice_accumulator(TmemStorage tmem_storage, int stage) { + return cute::make_tuple(tmem_storage.accumulators(_,_,_,stage)); + } + + template + CUTLASS_DEVICE static + auto + init_tmem_tensors(EpilogueTile epi_tile) { + TiledMma tiled_mma; + auto acc_shape = partition_accumulator_shape(); + // ((MMA_TILE_M,MMA_TILE_N),MMA_M,MMA_N,ACC_PIPE) — double-buffered for mainloop/epilogue overlap + Tensor accumulators = cutlass::detail::make_sm100_accumulator( + tiled_mma, acc_shape, EpilogueTile{}); + TmemStorage tmem_storage; + tmem_storage.accumulators = accumulators; + return tmem_storage; + } + + template + CUTLASS_DEVICE static + void + set_tmem_offsets(TmemStorage& tmem_storage, uint32_t tmem_base_addr) { + tmem_storage.accumulators.data() = tmem_base_addr; + } + + /// Set up the data needed by this collective for load. + template + CUTLASS_DEVICE auto + load_init( + ProblemShape_MNKL const& problem_shape_MNKL, + TensorStorage& shared_tensors) const { + using X = Underscore; + + auto [M,N,K,L] = problem_shape_MNKL; + + Tensor mA_mkl = observed_tma_load_a_->get_tma_tensor(make_shape(M,K,L)); + Tensor mB_nkl = observed_tma_load_b_->get_tma_tensor(make_shape(N,K,L)); + + Tensor gA_mkl = local_tile(mA_mkl, TileShape{}, make_coord(_,_,_), Step<_1, X,_1>{}); // (BLK_M, BLK_K, m, k, l) + Tensor gB_nkl = local_tile(mB_nkl, TileShape{}, make_coord(_,_,_), Step< X,_1,_1>{}); // (BLK_N, BLK_K, n, k, l) + + // Partition for this CTA + ThrMMA cta_mma = TiledMma{}.get_slice(blockIdx.x % size(typename TiledMma::AtomThrID{})); + + Tensor tCgA_mkl = cta_mma.partition_A(gA_mkl); // (MMA, MMA_M, MMA_K, m, k, l) + Tensor tCgB_nkl = cta_mma.partition_B(gB_nkl); // (MMA, MMA_N, MMA_K, n, k, l) + + Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); // (MMA,MMA_M,MMA_K,PIPE) + Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); // (MMA,MMA_N,MMA_K,PIPE) + + Layout cta_layout_mnk = make_layout(cluster_shape_); + Layout cta_layout_vmnk = tiled_divide(cta_layout_mnk, make_tile(typename TiledMma::AtomThrID{})); + auto cta_coord_vmnk = cta_layout_vmnk.get_flat_coord(block_rank_in_cluster_); + + auto [tAgA_mkl, tAsA] = tma_partition(*observed_tma_load_a_, + get<2>(cta_coord_vmnk), make_layout(size<2>(cta_layout_vmnk)), + group_modes<0,3>(sA), group_modes<0,3>(tCgA_mkl)); + + auto [tBgB_nkl, tBsB] = tma_partition(*observed_tma_load_b_, + get<1>(cta_coord_vmnk), make_layout(size<1>(cta_layout_vmnk)), + group_modes<0,3>(sB), group_modes<0,3>(tCgB_nkl)); + + uint16_t mcast_mask_a = create_tma_multicast_mask<2>(cta_layout_vmnk, cta_coord_vmnk); + uint16_t mcast_mask_b = create_tma_multicast_mask<1>(cta_layout_vmnk, cta_coord_vmnk); + + return LoadParams{ + shape<3>(gA_mkl), + tAgA_mkl, tBgB_nkl, tAsA, tBsB, + mcast_mask_a, mcast_mask_b}; + } + + /// Set up the data needed by this collective for mma compute. + template + CUTLASS_DEVICE auto + mma_init( + [[maybe_unused]] TmemStorage tmem_storage, + TensorStorage& shared_tensors) const { + + Tensor sA = make_tensor(make_smem_ptr(shared_tensors.smem_A.begin()), SmemLayoutA{}); + Tensor sB = make_tensor(make_smem_ptr(shared_tensors.smem_B.begin()), SmemLayoutB{}); + + Tensor tCrA = TiledMma::make_fragment_A(sA); // (MMA,MMA_M,MMA_K,PIPE) + Tensor tCrB = TiledMma::make_fragment_B(sB); // (MMA,MMA_N,MMA_K,PIPE) + + CUTE_STATIC_ASSERT_V(Int{} == size<3>(sA)); + CUTE_STATIC_ASSERT_V(Int{} == size<3>(sB)); + + TiledMma tiled_mma; + + if constexpr (IsRuntimeDataType) { + tiled_mma.idesc_.a_format_ = uint8_t(runtime_data_type_a_) & 0b111; + tiled_mma.idesc_.b_format_ = uint8_t(runtime_data_type_b_) & 0b111; + } + + return MmaParams{tiled_mma, tCrA, tCrB}; + } + + /// Perform a collective-scoped matrix multiply-accumulate + /// Producer Perspective + template < + class LoadParams, + class TileCoordMNKL, + class KTileIterator + > + CUTLASS_DEVICE auto + load( + MainloopPipeline mainloop_pipeline, + MainloopPipelineState mainloop_pipe_producer_state, + LoadParams const& load_inputs, + TileCoordMNKL const& cta_coord_mnkl, + KTileIterator k_tile_iter, int k_tile_count) { + + auto [unused_k_tiles, + tAgA_mkl, tBgB_nkl, tAsA, tBsB, + mcast_mask_a, mcast_mask_b] = load_inputs; + + Tensor tAgA = tAgA_mkl(_, get<0>(cta_coord_mnkl) / size(typename TiledMma::AtomThrID{}), _, get<3>(cta_coord_mnkl)); + Tensor tBgB = tBgB_nkl(_, get<1>(cta_coord_mnkl), _, get<3>(cta_coord_mnkl)); + + auto barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state); + + CUTLASS_PRAGMA_NO_UNROLL + while (k_tile_count > 0) { + mainloop_pipeline.producer_acquire(mainloop_pipe_producer_state, barrier_token); + + using BarrierType = typename MainloopPipeline::ProducerBarrierType; + BarrierType* tma_barrier = mainloop_pipeline.producer_get_barrier(mainloop_pipe_producer_state); + + int write_stage = mainloop_pipe_producer_state.index(); + ++mainloop_pipe_producer_state; + barrier_token = mainloop_pipeline.producer_try_acquire(mainloop_pipe_producer_state); + + if (cute::elect_one_sync()) { + copy(observed_tma_load_a_->with(*tma_barrier, mcast_mask_a), tAgA(_,*k_tile_iter), tAsA(_,write_stage)); + copy(observed_tma_load_b_->with(*tma_barrier, mcast_mask_b), tBgB(_,*k_tile_iter), tBsB(_,write_stage)); + } + + --k_tile_count; + ++k_tile_iter; + } + + return cute::make_tuple(mainloop_pipe_producer_state, k_tile_iter); + } + + /// Perform a Producer Epilogue to prevent early exit of ctas in a Cluster + CUTLASS_DEVICE void + load_tail(MainloopPipeline mainloop_pipeline, MainloopPipelineState mainloop_pipe_producer_state) { + mainloop_pipeline.producer_tail(mainloop_pipe_producer_state); + } + + /// Perform a collective-scoped matrix multiply-accumulate + /// Consumer Perspective + template < + class AccumulatorPipeline, + class FrgEngine, class FrgLayout, + class MmaParams, + class CtaTileCoord + > + CUTLASS_DEVICE auto + mma(cute::tuple pipelines, + cute::tuple pipeline_states, + cute::tuple> const& accumulators_pair, + MmaParams const& mma_inputs, + CtaTileCoord cta_tile_coord, + int k_tile_count + ) { + static_assert(is_tmem::value, "Accumulator must be tmem resident."); + static_assert(rank(FrgLayout{}) == 3, "Accumulator must be MMA-partitioned: (MMA, MMA_M, MMA_N)"); + + auto accumulators = get<0>(accumulators_pair); + auto [tiled_mma, tCrA, tCrB] = mma_inputs; + + auto [mainloop_pipeline, accumulator_pipeline] = pipelines; + auto [mainloop_pipe_consumer_state, accumulator_pipe_producer_state] = pipeline_states; + + uint32_t skip_wait = k_tile_count <= 0; + auto barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait); + + tiled_mma.accumulate_ = UMMA::ScaleOut::Zero; + + accumulator_pipeline.producer_acquire(accumulator_pipe_producer_state); + + CUTLASS_PRAGMA_NO_UNROLL + while (k_tile_count > 0) { + mainloop_pipeline.consumer_wait(mainloop_pipe_consumer_state, barrier_token); + + int read_stage = mainloop_pipe_consumer_state.index(); + auto curr_mainloop_pipe_consumer_state = mainloop_pipe_consumer_state; + + ++mainloop_pipe_consumer_state; + --k_tile_count; + skip_wait = k_tile_count <= 0; + barrier_token = mainloop_pipeline.consumer_try_wait(mainloop_pipe_consumer_state, skip_wait); + + CUTLASS_PRAGMA_UNROLL + for (int k_block = 0; k_block < size<2>(tCrA); ++k_block) { + if constexpr (WithBreuse) { + static_assert(size<1>(tCrA) == 2, + "The b-reuse feature expects size<1>(tCrA) == 2."); + static_assert(size<1>(tCrB) == 1, + "The b-reuse feature expects size<1>(tCrB) == 1."); + static_assert(size<1>(accumulators) == 2, + "The b-reuse feature expects size<1>(accumulators) == 2."); + + cute::gemm(tiled_mma.with(C{}), + tCrA(_,0,k_block,read_stage), + tCrB(_,0,k_block,read_stage), + accumulators(_,0,0)); + cute::gemm(tiled_mma.with(C{}), + tCrA(_,1,k_block,read_stage), + tCrB(_,0,k_block,read_stage), + accumulators(_,1,0)); + } + else { + cute::gemm(tiled_mma, + tCrA(_,_,k_block,read_stage), + tCrB(_,_,k_block,read_stage), + accumulators); + } + tiled_mma.accumulate_ = UMMA::ScaleOut::One; + } + mainloop_pipeline.consumer_release(curr_mainloop_pipe_consumer_state); + } + + return mainloop_pipe_consumer_state; + } + +protected: + + typename Params::TMA_A const* observed_tma_load_a_{nullptr}; + typename Params::TMA_B const* observed_tma_load_b_{nullptr}; + RuntimeDataTypeA runtime_data_type_a_{}; + RuntimeDataTypeB runtime_data_type_b_{}; + + ClusterShape cluster_shape_; + uint32_t block_rank_in_cluster_; +}; + +///////////////////////////////////////////////////////////////////////////////////////////////// + +} // namespace cutlass::gemm::collective + +///////////////////////////////////////////////////////////////////////////////////////////////// diff --git a/include/cutlass/gemm/dispatch_policy.hpp b/include/cutlass/gemm/dispatch_policy.hpp index 9cd9d25691..969d9e141a 100644 --- a/include/cutlass/gemm/dispatch_policy.hpp +++ b/include/cutlass/gemm/dispatch_policy.hpp @@ -37,6 +37,7 @@ #include "cute/numeric/integral_constant.hpp" // cute::false_type #include "cute/atom/copy_traits_sm100.hpp" #include "cutlass/detail/collective/sm103_kernel_type.hpp" + ////////////////////////////////////////////////////////////////////////////// namespace cutlass::detail { @@ -579,6 +580,24 @@ struct KernelPtrArrayTmaWarpSpecializedInputTransformSm100 final { static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_; }; +// SM107 kernel schedules +template< + int SchedulerPipelineStageCount_, + int AccumulatorPipelineStageCount_ +> +struct KernelTmaWarpSpecializedSm107 final { + static constexpr int SchedulerPipelineStageCount = SchedulerPipelineStageCount_; + static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_; +}; + +template< + int SchedulerPipelineStageCount_, + int AccumulatorPipelineStageCount_ +> +struct KernelTmaWarpSpecializedBlockScaledSm107 final { + static constexpr int SchedulerPipelineStageCount = SchedulerPipelineStageCount_; + static constexpr int AccumulatorPipelineStageCount = AccumulatorPipelineStageCount_; +}; // SM120 kernel schedules template @@ -912,6 +931,65 @@ using KernelPtrArrayTmaWarpSpecialized2SmBlockScaledMxNvf4UltraVs16Sm103 = Kerne using KernelPtrArrayTmaWarpSpecialized1SmBlockScaledMxNvf4UltraVs32Sm103 = KernelPtrArrayTmaWarpSpecialized1SmBlockScaledMxNvf4UltraVs32Sm103DisablePrefetch; using KernelPtrArrayTmaWarpSpecialized2SmBlockScaledMxNvf4UltraVs32Sm103 = KernelPtrArrayTmaWarpSpecialized2SmBlockScaledMxNvf4UltraVs32Sm103DisablePrefetch; +/////////////////////////////////////////////////////////////////////////////////////////////////////// +// +// SM107 Dispatch Policies +// +/////////////////////////////////////////////////////////////////////////////////////////////////////// + +struct KernelScheduleSm107 {}; + +/////////////////////////////////////////////////////////////////////////////////////////////////////// +// SM107 Dense GEMM Dispatch Policies +/////////////////////////////////////////////////////////////////////////////////////////////////////// + +struct KernelScheduleSm107DenseGemm : KernelScheduleSm107 {}; +struct KernelScheduleSm107DenseGemmf8f6f4 : KernelScheduleSm107DenseGemm {}; +struct KernelScheduleSm107BlockScaledMxf8f6f4 : KernelScheduleSm107DenseGemm {}; +struct KernelScheduleSm107BlockScaledMxNvf4 : KernelScheduleSm107DenseGemm {}; + +struct KernelScheduleSm107DenseGemmf8f6f4WithoutBreuse : KernelScheduleSm107DenseGemmf8f6f4 {}; +struct KernelScheduleSm107DenseGemmf8f6f4WithBreuse : KernelScheduleSm107DenseGemmf8f6f4 {}; + +struct KernelScheduleSm107BlockScaledMxf8f6f4WithoutBreuse : KernelScheduleSm107BlockScaledMxf8f6f4 {}; +struct KernelScheduleSm107BlockScaledMxf8f6f4WithBreuse : KernelScheduleSm107BlockScaledMxf8f6f4 {}; + +// Block Scaled GEMM (NVF4): Specialize for scale factor vector size, then B-reuse +struct KernelScheduleSm107BlockScaledMxNvf4Vs16 : KernelScheduleSm107BlockScaledMxNvf4 {}; +struct KernelScheduleSm107BlockScaledMxNvf4Vs32 : KernelScheduleSm107BlockScaledMxNvf4 {}; + +struct KernelScheduleSm107BlockScaledMxNvf4Vs16WithoutBreuse : KernelScheduleSm107BlockScaledMxNvf4Vs16 {}; +struct KernelScheduleSm107BlockScaledMxNvf4Vs16WithBreuse : KernelScheduleSm107BlockScaledMxNvf4Vs16 {}; +struct KernelScheduleSm107BlockScaledMxNvf4Vs32WithoutBreuse : KernelScheduleSm107BlockScaledMxNvf4Vs32 {}; +struct KernelScheduleSm107BlockScaledMxNvf4Vs32WithBreuse : KernelScheduleSm107BlockScaledMxNvf4Vs32 {}; + +// Dense GEMM: Specialize for 1SM vs 2SM +struct KernelTmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithoutBreuse final : KernelSchedule1Sm, KernelScheduleSm107DenseGemmf8f6f4WithoutBreuse {}; +struct KernelTmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithoutBreuse final : KernelSchedule2Sm, KernelScheduleSm107DenseGemmf8f6f4WithoutBreuse {}; + +struct KernelTmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithBreuse final : KernelSchedule1Sm, KernelScheduleSm107DenseGemmf8f6f4WithBreuse {}; +struct KernelTmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithBreuse final : KernelSchedule2Sm, KernelScheduleSm107DenseGemmf8f6f4WithBreuse {}; + +// Dense blockscaled GEMM: specialized for 1SM vs 2SM +struct KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse final : KernelSchedule1Sm, KernelScheduleSm107BlockScaledMxf8f6f4WithoutBreuse {}; +struct KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse final : KernelSchedule2Sm, KernelScheduleSm107BlockScaledMxf8f6f4WithoutBreuse {}; + +struct KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse final : KernelSchedule1Sm, KernelScheduleSm107BlockScaledMxf8f6f4WithBreuse {}; +struct KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse final : KernelSchedule2Sm, KernelScheduleSm107BlockScaledMxf8f6f4WithBreuse {}; + +// Dense blockscaled GEMM (NVF4): specialized for 1SM vs 2SM, scale factor vector size, and B-reuse +struct KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse final : KernelSchedule1Sm, KernelScheduleSm107BlockScaledMxNvf4Vs16WithoutBreuse {}; +struct KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse final : KernelSchedule2Sm, KernelScheduleSm107BlockScaledMxNvf4Vs16WithoutBreuse {}; + +struct KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse final : KernelSchedule1Sm, KernelScheduleSm107BlockScaledMxNvf4Vs16WithBreuse {}; +struct KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse final : KernelSchedule2Sm, KernelScheduleSm107BlockScaledMxNvf4Vs16WithBreuse {}; + +struct KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse final : KernelSchedule1Sm, KernelScheduleSm107BlockScaledMxNvf4Vs32WithoutBreuse {}; +struct KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse final : KernelSchedule2Sm, KernelScheduleSm107BlockScaledMxNvf4Vs32WithoutBreuse {}; + +struct KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse final : KernelSchedule1Sm, KernelScheduleSm107BlockScaledMxNvf4Vs32WithBreuse {}; +struct KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse final : KernelSchedule2Sm, KernelScheduleSm107BlockScaledMxNvf4Vs32WithBreuse {}; + /////////////////////////////////////////////////////////////////////////////////////////////////////// // // SM120 Dispatch Policies @@ -1445,6 +1523,39 @@ struct MainloopSm103ArrayTmaUmmaWarpSpecializedBlockScaled { constexpr static cutlass::sm103::detail::KernelPrefetchType PrefetchType = PrefetchType_; }; +// n-buffer in smem, pipelined with Rubin UMMA and TMA, Warp specialized dynamic schedule +template< + int Stages_, + int SchedulerPipelineStageCount_, + int AccumulatorPipelineStageCount_, + class ClusterShape_ = Shape<_1,_1,_1>, + bool WithBreuse_ = false, + class ArchTag_ = arch::Sm107 +> +struct MainloopSm107TmaUmmaWarpSpecialized { + constexpr static int Stages = Stages_; + constexpr static bool WithBreuse = WithBreuse_; + using ClusterShape = ClusterShape_; + using ArchTag = ArchTag_; + using Schedule = KernelTmaWarpSpecializedSm107; +}; + +template< + int Stages_, + int SchedulerPipelineStageCount_, + int AccumulatorPipelineStageCount_, + class ClusterShape_ = Shape<_1,_1,_1>, + bool WithBreuse_ = false, + class ArchTag_ = arch::Sm107 +> +struct MainloopSm107TmaUmmaWarpSpecializedBlockScaled { + constexpr static int Stages = Stages_; + constexpr static bool WithBreuse = WithBreuse_; + using ClusterShape = ClusterShape_; + using ArchTag = ArchTag_; + using Schedule = KernelTmaWarpSpecializedBlockScaledSm107; +}; + template< int Stages_, int SchedulerPipelineStageCount_, diff --git a/include/cutlass/gemm/kernel/gemm_universal.hpp b/include/cutlass/gemm/kernel/gemm_universal.hpp index d0c84d3e9d..7e5abb44f8 100644 --- a/include/cutlass/gemm/kernel/gemm_universal.hpp +++ b/include/cutlass/gemm/kernel/gemm_universal.hpp @@ -76,6 +76,7 @@ struct IsCutlass3ArrayKernel +class GemmUniversal< + ProblemShape_, + CollectiveMainloop_, + CollectiveEpilogue_, + TileSchedulerTag_, + cute::enable_if_t< + cute::disjunction_v, + cutlass::detail::is_kernel_tag_of>>> +{ +public: + // + // Type Aliases + // + using ProblemShape = ProblemShape_; + static_assert(rank(ProblemShape{}) == 3 or rank(ProblemShape{}) == 4, + "ProblemShape{} should be or "); + + // Mainloop derived types + using CollectiveMainloop = CollectiveMainloop_; + using TileShape = typename CollectiveMainloop::TileShape; + using TiledMma = typename CollectiveMainloop::TiledMma; + using ArchTag = typename CollectiveMainloop::ArchTag; + using ElementA = typename CollectiveMainloop::ElementA; + using StrideA = typename CollectiveMainloop::StrideA; + using ElementB = typename CollectiveMainloop::ElementB; + using StrideB = typename CollectiveMainloop::StrideB; + using LayoutSFA = typename cutlass::detail::LayoutSFAType::type; + using LayoutSFB = typename cutlass::detail::LayoutSFBType::type; + using ElementSF = typename cutlass::detail::ElementSFType::type; + using DispatchPolicy = typename CollectiveMainloop::DispatchPolicy; + using ElementAccumulator = typename CollectiveMainloop::ElementAccumulator; + using ClusterShape = typename DispatchPolicy::ClusterShape; + using MainloopArguments = typename CollectiveMainloop::Arguments; + using MainloopParams = typename CollectiveMainloop::Params; + static_assert(ArchTag::kMinComputeCapability >= 107); + + // Epilogue derived types + using CollectiveEpilogue = CollectiveEpilogue_; + using EpilogueTile = typename CollectiveEpilogue::EpilogueTile; + using ElementC = typename CollectiveEpilogue::ElementC; + using StrideC = typename CollectiveEpilogue::StrideC; + using ElementD = typename CollectiveEpilogue::ElementD; + using StrideD = typename CollectiveEpilogue::StrideD; + using EpilogueArguments = typename CollectiveEpilogue::Arguments; + using EpilogueParams = typename CollectiveEpilogue::Params; + static constexpr bool IsComplex = CollectiveEpilogue::NumAccumulatorMtxs == 2; + + // CLC pipeline depth + // determines how many waves (stages-1) a warp can race ahead + static constexpr uint32_t SchedulerPipelineStageCount = DispatchPolicy::Schedule::SchedulerPipelineStageCount; + static constexpr uint32_t AccumulatorPipelineStageCount = DispatchPolicy::Schedule::AccumulatorPipelineStageCount; + + // TileID scheduler + // Get Blk and Scheduling tile shapes + using AtomThrShapeMNK = typename CollectiveMainloop::AtomThrShapeMNK; + using CtaShape_MNK = typename CollectiveMainloop::CtaShape_MNK; + using TileSchedulerTag = TileSchedulerTag_; + using TileScheduler = typename detail::TileSchedulerSelector< + TileSchedulerTag, ArchTag, CtaShape_MNK, ClusterShape, SchedulerPipelineStageCount>::Scheduler; + using TileSchedulerArguments = typename TileScheduler::Arguments; + using TileSchedulerParams = typename TileScheduler::Params; + + static constexpr bool IsSchedDynamicPersistent = TileScheduler::IsDynamicPersistent; + static constexpr bool IsDynamicCluster = not cute::is_static_v; + static constexpr bool IsGdcEnabled = cutlass::arch::IsGdcGloballyEnabled; + + // Warp specialization thread count per threadblock + static constexpr uint32_t NumSchedThreads = NumThreadsPerWarp; // 1 warp + static constexpr uint32_t NumMMAThreads = NumThreadsPerWarp; // 1 warp + static constexpr uint32_t NumMainloopLoadThreads = NumThreadsPerWarp; // 1 warp + static constexpr uint32_t NumEpilogueLoadThreads = NumThreadsPerWarp; // 1 warp + static constexpr uint32_t NumEpilogueThreads = CollectiveEpilogue::ThreadCount; + static constexpr uint32_t NumEpilogueWarps = NumEpilogueThreads / NumThreadsPerWarp; + + static constexpr uint32_t MaxThreadsPerBlock = NumSchedThreads + + NumMainloopLoadThreads + NumMMAThreads + + NumEpilogueLoadThreads + NumEpilogueThreads; + static constexpr uint32_t MinBlocksPerMultiprocessor = 1; + + static constexpr uint32_t NumEpilogueSubTiles = CollectiveEpilogue::get_load_pipe_increment(CtaShape_MNK{}); + + // Fixup performed for split-/stream-K is done across warps in different CTAs + // at epilogue subtile granularity. Thus, there must be one barrier per sub-tile per + // epilogue warp. + static constexpr uint32_t NumFixupBarriers = 1; + static constexpr uint32_t CLCResponseSize = sizeof(typename TileScheduler::CLCResponse); + + // Pipeline and pipeline state types + using MainloopPipeline = typename CollectiveMainloop::MainloopPipeline; + using MainloopPipelineState = typename CollectiveMainloop::MainloopPipelineState; + using EpiLoadPipeline = typename CollectiveEpilogue::LoadPipeline; + using EpiLoadPipelineState = typename CollectiveEpilogue::LoadPipelineState; + + using EpiStorePipeline = typename CollectiveEpilogue::StorePipeline; + using EpiStorePipelineState = typename CollectiveEpilogue::StorePipelineState; + + using LoadOrderBarrier = cutlass::OrderedSequenceBarrier<1,2>; + + using AccumulatorPipeline = cutlass::PipelineUmmaAsync; + using AccumulatorPipelineState = typename AccumulatorPipeline::PipelineState; + + using CLCPipeline = cutlass::PipelineCLCFetchAsync; + using CLCPipelineState = typename CLCPipeline::PipelineState; + + using CLCThrottlePipeline = cutlass::PipelineAsync; + using CLCThrottlePipelineState = typename CLCThrottlePipeline::PipelineState; + + using TmemAllocator = cute::conditional_t(typename TiledMma::ThrLayoutVMNK{})) == 1, + cute::TMEM::Allocator1Sm, cute::TMEM::Allocator2Sm>; + + // Kernel level shared memory storage + struct SharedStorage { + struct PipelineStorage : cute::aligned_struct<16, _1> { + using MainloopPipelineStorage = typename CollectiveMainloop::PipelineStorage; + using EpiLoadPipelineStorage = typename CollectiveEpilogue::PipelineStorage; + using LoadOrderBarrierStorage = typename LoadOrderBarrier::SharedStorage; + using CLCPipelineStorage = typename CLCPipeline::SharedStorage; + using AccumulatorPipelineStorage = typename AccumulatorPipeline::SharedStorage; + using CLCThrottlePipelineStorage = typename CLCThrottlePipeline::SharedStorage; + + alignas(16) MainloopPipelineStorage mainloop; + alignas(16) EpiLoadPipelineStorage epi_load; + alignas(16) LoadOrderBarrierStorage load_order; + alignas(16) CLCPipelineStorage clc; + alignas(16) AccumulatorPipelineStorage accumulator; + alignas(16) CLCThrottlePipelineStorage clc_throttle; + alignas(16) arch::ClusterBarrier tmem_dealloc; + } pipelines; + + alignas(16) typename TileScheduler::CLCResponse clc_response[SchedulerPipelineStageCount]; + uint32_t tmem_base_ptr; + + struct TensorStorage : cute::aligned_struct<128, _1> { + using EpilogueTensorStorage = typename CollectiveEpilogue::TensorStorage; + using MainloopTensorStorage = typename CollectiveMainloop::TensorStorage; + + EpilogueTensorStorage epilogue; + MainloopTensorStorage mainloop; + } tensors; + }; + + static constexpr int SharedStorageSize = sizeof(SharedStorage); + + // Host facing host arguments + struct Arguments { + GemmUniversalMode mode{}; + ProblemShape problem_shape{}; + MainloopArguments mainloop{}; + EpilogueArguments epilogue{}; + KernelHardwareInfo hw_info{}; + TileSchedulerArguments scheduler{}; + }; + + // Kernel device entry point API + struct Params { + GemmUniversalMode mode{}; + ProblemShape problem_shape{}; + MainloopParams mainloop{}; + EpilogueParams epilogue{}; + TileSchedulerParams scheduler{}; + KernelHardwareInfo hw_info{}; + }; + + enum class WarpCategory : int32_t { + MMA = 0, + Sched = 1, + MainloopLoad = 2, + EpilogueLoad = 3, + Epilogue = 4 + }; + + struct IsParticipant { + uint32_t mma = false; + uint32_t sched = false; + uint32_t main_load = false; + uint32_t epi_load = false; + uint32_t epilogue = false; + }; + + // + // Methods + // + + // Convert to underlying arguments. + static + Params + to_underlying_arguments(Arguments const& args, void* workspace) { + (void) workspace; + auto problem_shape = args.problem_shape; + auto problem_shape_MNKL = append<4>(problem_shape, 1); + + // Get SM count if needed, otherwise use user supplied SM count + int sm_count = args.hw_info.sm_count; + if (sm_count != 0) { + CUTLASS_TRACE_HOST(" WARNING: SM107 tile scheduler does not allow for user specified SM counts.\n" + " To restrict a kernel's resource usage, consider using CUDA driver APIs instead (green contexts)."); + } + CUTLASS_TRACE_HOST("to_underlying_arguments(): Setting persistent grid SM count to " << sm_count); + + // Calculate workspace pointers + uint8_t* workspace_ptr = reinterpret_cast(workspace); + size_t workspace_offset = 0; + + // Epilogue + void* epilogue_workspace = workspace_ptr + workspace_offset; + workspace_offset += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue); + workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment); + + void* mainloop_workspace = nullptr; + + // Tile scheduler + void* scheduler_workspace = workspace_ptr + workspace_offset; + workspace_offset += TileScheduler::template get_workspace_size( + args.scheduler, args.problem_shape, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs); + workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment); + + return { + args.mode, + args.problem_shape, + CollectiveMainloop::to_underlying_arguments(args.problem_shape, args.mainloop, mainloop_workspace, args.hw_info), + CollectiveEpilogue::to_underlying_arguments(args.problem_shape, args.epilogue, epilogue_workspace), + TileScheduler::to_underlying_arguments( + problem_shape_MNKL, TileShape{}, AtomThrShapeMNK{}, ClusterShape{}, + args.hw_info, args.scheduler, scheduler_workspace + ) + ,args.hw_info + }; + } + + static bool + can_implement(Arguments const& args) { + bool implementable = (args.mode == GemmUniversalMode::kGemm) or + (args.mode == GemmUniversalMode::kBatched && rank(ProblemShape{}) == 4); + if (!implementable) { + CUTLASS_TRACE_HOST(" CAN IMPLEMENT: Arguments or Problem Shape don't meet the requirements.\n"); + return implementable; + } + implementable &= CollectiveMainloop::can_implement(args.problem_shape, args.mainloop); + implementable &= CollectiveEpilogue::can_implement(args.problem_shape, args.epilogue); + implementable &= TileScheduler::can_implement(args.scheduler, args.hw_info); + + if constexpr (IsDynamicCluster) { + static constexpr int MaxClusterSize = 16; + implementable &= size(args.hw_info.cluster_shape) <= MaxClusterSize; + implementable &= size(args.hw_info.cluster_shape_fallback) <= MaxClusterSize; + implementable &= cutlass::detail::preferred_cluster_can_implement(args.hw_info.cluster_shape, args.hw_info.cluster_shape_fallback); + } + + constexpr bool IsBlockscaled = !cute::is_void_v; + if constexpr (IsBlockscaled) { + if constexpr (IsDynamicCluster) { + implementable &= cutlass::detail::preferred_cluster_can_implement(args.hw_info.cluster_shape, args.hw_info.cluster_shape_fallback); + // Special cluster shape check for scale factor multicasts. Due to limited size of scale factors, we can't multicast among + // more than 4 CTAs + implementable &= (args.hw_info.cluster_shape.x <= 4 && args.hw_info.cluster_shape.y <= 4 && + args.hw_info.cluster_shape_fallback.x <= 4 && args.hw_info.cluster_shape_fallback.y <= 4); + } + else { + // Special cluster shape check for scale factor multicasts. Due to limited size of scale factors, we can't multicast among + // more than 4 CTAs + implementable &= ((size<0>(ClusterShape{}) <= 4) && (size<1>(ClusterShape{}) <= 4)); + } + } + + return implementable; + } + + static size_t + get_workspace_size(Arguments const& args) { + size_t workspace_size = 0; + + // Epilogue + workspace_size += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue); + workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment); + + // Tile scheduler + workspace_size += TileScheduler::template get_workspace_size( + args.scheduler, args.problem_shape, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs); + workspace_size = round_nearest(workspace_size, MinWorkspaceAlignment); + + return workspace_size; + } + + static cutlass::Status + initialize_workspace(Arguments const& args, void* workspace = nullptr, cudaStream_t stream = nullptr, + CudaHostAdapter* cuda_adapter = nullptr) { + Status status = Status::kSuccess; + uint8_t* workspace_ptr = reinterpret_cast(workspace); + size_t workspace_offset = 0; + + // Epilogue + status = CollectiveEpilogue::initialize_workspace(args.problem_shape, args.epilogue, workspace_ptr + workspace_offset, stream, cuda_adapter); + workspace_offset += CollectiveEpilogue::get_workspace_size(args.problem_shape, args.epilogue); + workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment); + if (status != Status::kSuccess) { + return status; + } + + // Tile scheduler + status = TileScheduler::template initialize_workspace( + args.scheduler, workspace_ptr + workspace_offset, stream, args.problem_shape, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs, cuda_adapter); + workspace_offset += TileScheduler::template get_workspace_size( + args.scheduler, args.problem_shape, args.hw_info, NumFixupBarriers, NumEpilogueSubTiles, CollectiveEpilogue::NumAccumulatorMtxs); + workspace_offset = round_nearest(workspace_offset, MinWorkspaceAlignment); + if (status != Status::kSuccess) { + return status; + } + + return status; + } + + // Computes the kernel launch grid shape based on runtime parameters + static dim3 + get_grid_shape(Params const& params) { + // NOTE cluster_shape here is the major cluster shape, not fallback one + auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}, params.hw_info.cluster_shape); + + auto problem_shape_MNKL = append<4>(params.problem_shape, Int<1>{}); + return TileScheduler::get_grid_shape( + params.scheduler, + problem_shape_MNKL, + TileShape{}, + AtomThrShapeMNK{}, + cluster_shape, + params.hw_info); + } + + static dim3 + get_block_shape() { + return dim3(MaxThreadsPerBlock, 1, 1); + } + + CUTLASS_DEVICE + void + operator() (Params const& params, char* smem_buf) { + + using namespace cute; + using X = Underscore; + + static_assert(SharedStorageSize <= cutlass::arch::sm107_smem_capacity_bytes, "SMEM usage exceeded capacity."); + + // Separate out problem shape for convenience + // Optionally append 1s until problem shape is rank-4 in case its is only rank-3 (MNK) + auto problem_shape_MNKL = append<4>(params.problem_shape, Int<1>{}); + auto [M,N,K,L] = problem_shape_MNKL; + + // Account for more than one epilogue warp + int warp_idx = canonical_warp_idx_sync(); + WarpCategory warp_category = warp_idx < static_cast(WarpCategory::Epilogue) ? WarpCategory(warp_idx) + : WarpCategory::Epilogue; + uint32_t lane_predicate = cute::elect_one_sync(); + auto cluster_shape = cutlass::detail::select_cluster_shape(ClusterShape{}); + int cluster_size = size(cluster_shape); + uint32_t cta_rank_in_cluster = cute::block_rank_in_cluster(); + bool is_first_cta_in_cluster = cta_rank_in_cluster == 0; + int cta_coord_v = cta_rank_in_cluster % size<0>(typename TiledMma::AtomThrID{}); + bool is_mma_leader_cta = cta_coord_v == 0; + constexpr bool has_mma_peer_cta = size(AtomThrShapeMNK{}) == 2; + [[maybe_unused]] uint32_t mma_peer_cta_rank = has_mma_peer_cta ? cta_rank_in_cluster ^ 1 : cta_rank_in_cluster; + + // Kernel level shared memory storage + SharedStorage& shared_storage = *reinterpret_cast(smem_buf); + + // In a warp specialized kernel, collectives expose data movement and compute operations separately + CollectiveMainloop collective_mainloop(params.mainloop, cluster_shape, cta_rank_in_cluster); + CollectiveEpilogue collective_epilogue(params.epilogue, shared_storage.tensors.epilogue); + + // Issue Tma Descriptor Prefetch from a single thread + if ((warp_category == WarpCategory::Sched) && lane_predicate) { + collective_mainloop.prefetch_tma_descriptors(); + } + if ((warp_category == WarpCategory::EpilogueLoad) && lane_predicate) { + collective_epilogue.prefetch_tma_descriptors(params.epilogue); + } + + bool is_epi_load_needed = collective_epilogue.is_producer_load_needed(); + IsParticipant is_participant = { + (warp_category == WarpCategory::MMA), // mma + (warp_category == WarpCategory::Sched) && is_first_cta_in_cluster, // sched + (warp_category == WarpCategory::MainloopLoad), // main_load + (warp_category == WarpCategory::EpilogueLoad) && is_epi_load_needed, // epi_load + (warp_category == WarpCategory::Epilogue) // epilogue + }; + + // Mainloop Load pipeline + typename MainloopPipeline::Params mainloop_pipeline_params; + if (WarpCategory::MainloopLoad == warp_category) { + mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Producer; + } + if (WarpCategory::MMA == warp_category) { + mainloop_pipeline_params.role = MainloopPipeline::ThreadCategory::Consumer; + } + mainloop_pipeline_params.is_leader = lane_predicate && is_mma_leader_cta && is_participant.main_load; + mainloop_pipeline_params.transaction_bytes = CollectiveMainloop::TmaTransactionBytes; + mainloop_pipeline_params.initializing_warp = 0; + MainloopPipeline mainloop_pipeline(shared_storage.pipelines.mainloop, + mainloop_pipeline_params, + cluster_shape, + cute::true_type{}, // Perform barrier init + cute::false_type{}); // Delay mask calculation + + // Epilogue Load pipeline + typename EpiLoadPipeline::Params epi_load_pipeline_params; + if (WarpCategory::EpilogueLoad == warp_category) { + epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Producer; + } + if (WarpCategory::Epilogue == warp_category) { + epi_load_pipeline_params.role = EpiLoadPipeline::ThreadCategory::Consumer; + } + epi_load_pipeline_params.dst_blockid = cta_rank_in_cluster; + epi_load_pipeline_params.producer_arv_count = NumEpilogueLoadThreads; + epi_load_pipeline_params.consumer_arv_count = NumEpilogueThreads; + epi_load_pipeline_params.transaction_bytes = CollectiveEpilogue::TmaTransactionBytes; + epi_load_pipeline_params.initializing_warp = 1; + EpiLoadPipeline epi_load_pipeline(shared_storage.pipelines.epi_load, epi_load_pipeline_params); + + // Epilogue Store pipeline + typename EpiStorePipeline::Params epi_store_pipeline_params; + epi_store_pipeline_params.always_wait = true; + EpiStorePipeline epi_store_pipeline(epi_store_pipeline_params); + + // Load order barrier + typename LoadOrderBarrier::Params load_order_barrier_params; + load_order_barrier_params.group_id = (warp_category == WarpCategory::MainloopLoad) ? 0 : 1; + load_order_barrier_params.group_size = NumMainloopLoadThreads; + load_order_barrier_params.initializing_warp = 3; + LoadOrderBarrier load_order_barrier(shared_storage.pipelines.load_order, load_order_barrier_params); + + // CLC pipeline + typename CLCPipeline::Params clc_pipeline_params; + if (WarpCategory::Sched == warp_category) { + clc_pipeline_params.role = CLCPipeline::ThreadCategory::ProducerConsumer; + } + else { + clc_pipeline_params.role = CLCPipeline::ThreadCategory::Consumer; + } + clc_pipeline_params.producer_blockid = 0; + clc_pipeline_params.producer_arv_count = 1; + clc_pipeline_params.consumer_arv_count = NumSchedThreads + cluster_size * + (NumMainloopLoadThreads + NumEpilogueThreads + NumMMAThreads); + if (is_epi_load_needed) { + clc_pipeline_params.consumer_arv_count += cluster_size * NumEpilogueLoadThreads; + } + clc_pipeline_params.transaction_bytes = CLCResponseSize; + clc_pipeline_params.initializing_warp = 4; + CLCPipeline clc_pipeline(shared_storage.pipelines.clc, clc_pipeline_params, cluster_shape); + + // Mainloop-Epilogue pipeline + typename AccumulatorPipeline::Params accumulator_pipeline_params; + if (WarpCategory::MMA == warp_category) { + accumulator_pipeline_params.role = AccumulatorPipeline::ThreadCategory::Producer; + } + if (WarpCategory::Epilogue == warp_category) { + accumulator_pipeline_params.role = AccumulatorPipeline::ThreadCategory::Consumer; + } + // Only one producer thread arrives on this barrier. + accumulator_pipeline_params.producer_arv_count = 1; + accumulator_pipeline_params.consumer_arv_count = size(AtomThrShapeMNK{}) * NumEpilogueThreads; + accumulator_pipeline_params.initializing_warp = 5; + AccumulatorPipeline accumulator_pipeline(shared_storage.pipelines.accumulator, + accumulator_pipeline_params, + cluster_shape, + cute::true_type{}, // Perform barrier init + cute::false_type{}); // Delay mask calculation + + // CLC throttle pipeline + typename CLCThrottlePipeline::Params clc_throttle_pipeline_params; + if (WarpCategory::MainloopLoad == warp_category) { + clc_throttle_pipeline_params.role = CLCThrottlePipeline::ThreadCategory::Producer; + } + if (WarpCategory::Sched == warp_category) { + clc_throttle_pipeline_params.role = CLCThrottlePipeline::ThreadCategory::Consumer; + } + clc_throttle_pipeline_params.producer_arv_count = NumMainloopLoadThreads; + clc_throttle_pipeline_params.consumer_arv_count = NumSchedThreads; + clc_throttle_pipeline_params.dst_blockid = 0; + clc_throttle_pipeline_params.initializing_warp = 3; + CLCThrottlePipeline clc_throttle_pipeline(shared_storage.pipelines.clc_throttle, clc_throttle_pipeline_params); + CLCThrottlePipelineState clc_pipe_throttle_consumer_state; + CLCThrottlePipelineState clc_pipe_throttle_producer_state = cutlass::make_producer_start_state(); + + // Tmem allocator + TmemAllocator tmem_allocator{}; + + // Sync allocation status between MMA and epilogue warps within CTA + arch::NamedBarrier tmem_allocation_result_barrier(NumMMAThreads + NumEpilogueThreads, cutlass::arch::ReservedNamedBarriers::TmemAllocBarrier); + // Sync deallocation status between MMA warps of peer CTAs + arch::ClusterBarrier& tmem_deallocation_result_barrier = shared_storage.pipelines.tmem_dealloc; + [[maybe_unused]] uint32_t dealloc_barrier_phase = 0; + if (WarpCategory::MMA == warp_category) { + if (has_mma_peer_cta && lane_predicate) { + tmem_deallocation_result_barrier.init(NumMMAThreads); + } + } + + // We need this to guarantee that the Pipeline init is visible + // To all producers and consumer threadblocks in the cluster + pipeline_init_arrive_relaxed(cluster_size); + + auto load_inputs = collective_mainloop.load_init( + problem_shape_MNKL, shared_storage.tensors.mainloop); + + MainloopPipelineState mainloop_pipe_consumer_state; + MainloopPipelineState mainloop_pipe_producer_state = cutlass::make_producer_start_state(); + + EpiLoadPipelineState epi_load_pipe_consumer_state; + EpiLoadPipelineState epi_load_pipe_producer_state = cutlass::make_producer_start_state(); + + // epilogue store pipe is producer-only (consumer is TMA unit, waits via scoreboarding) + EpiStorePipelineState epi_store_pipe_producer_state = cutlass::make_producer_start_state(); + + CLCPipelineState clc_pipe_consumer_state; + CLCPipelineState clc_pipe_producer_state = cutlass::make_producer_start_state(); + + AccumulatorPipelineState accumulator_pipe_consumer_state; + AccumulatorPipelineState accumulator_pipe_producer_state = cutlass::make_producer_start_state(); + + dim3 block_id_in_cluster = cute::block_id_in_cluster(); + + // Calculate mask after cluster barrier arrival + mainloop_pipeline.init_masks(cluster_shape, block_id_in_cluster); + accumulator_pipeline.init_masks(cluster_shape, block_id_in_cluster); + + // TileID scheduler + TileScheduler scheduler(&shared_storage.clc_response[0], params.scheduler, block_id_in_cluster); + typename TileScheduler::WorkTileInfo work_tile_info = scheduler.initial_work_tile_info(cluster_shape); + auto cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info); + // + // TMEM "Allocation" + // + auto tmem_storage = collective_mainloop.template init_tmem_tensors(EpilogueTile{}); + + pipeline_init_wait(cluster_size); + + if (is_participant.main_load) { + // Ensure that the prefetched kernel does not touch + // unflushed global memory prior to this instruction + cutlass::arch::wait_on_dependent_grids(); + bool do_load_order_arrive = is_epi_load_needed; + bool requires_clc_query = true; + + do { + // Get the number of K tiles to compute for this work as well as the starting K tile offset of the work. + auto k_tile_iter = scheduler.get_k_tile_iterator(work_tile_info, problem_shape_MNKL, CtaShape_MNK{}, load_inputs.k_tiles); + auto k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, problem_shape_MNKL, CtaShape_MNK{}); + auto k_tile_prologue = min(MainloopPipeline::Stages, k_tile_count); + + if constexpr (IsSchedDynamicPersistent) { + if (is_first_cta_in_cluster && requires_clc_query) { + clc_throttle_pipeline.producer_acquire(clc_pipe_throttle_producer_state); + clc_throttle_pipeline.producer_commit(clc_pipe_throttle_producer_state); + ++clc_pipe_throttle_producer_state; + } + } + + // Start mainloop prologue loads, arrive on the epilogue residual load barrier, resume mainloop loads + auto [mainloop_producer_state_next, k_tile_iter_next] = collective_mainloop.load( + mainloop_pipeline, + mainloop_pipe_producer_state, + load_inputs, + cta_coord_mnkl, + k_tile_iter, k_tile_prologue + ); + mainloop_pipe_producer_state = mainloop_producer_state_next; + + if (do_load_order_arrive) { + load_order_barrier.arrive(); + do_load_order_arrive = false; + } + + auto [mainloop_producer_state_next_, unused_] = collective_mainloop.load( + mainloop_pipeline, + mainloop_pipe_producer_state, + load_inputs, + cta_coord_mnkl, + k_tile_iter_next, k_tile_count - k_tile_prologue + ); + mainloop_pipe_producer_state = mainloop_producer_state_next_; + + // Sync warp to prevent non-participating threads entering next wave early + __syncwarp(); + + auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work( + work_tile_info, + clc_pipeline, + clc_pipe_consumer_state + ); + work_tile_info = next_work_tile_info; + cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info); + requires_clc_query = increment_pipe; + if (increment_pipe) { + ++clc_pipe_consumer_state; + } + } while (work_tile_info.is_valid()); + collective_mainloop.load_tail(mainloop_pipeline, mainloop_pipe_producer_state); + + } + + else if (is_participant.sched) { + if constexpr (IsSchedDynamicPersistent) { + // Whether a new CLC query must be performed. + // See comment below where this variable is updated for a description of + // why this variable is needed. + bool requires_clc_query = true; + + cutlass::arch::wait_on_dependent_grids(); + do { + if (requires_clc_query) { + // Throttle CLC query to mitigate workload imbalance caused by skews among persistent workers. + clc_throttle_pipeline.consumer_wait(clc_pipe_throttle_consumer_state); + clc_throttle_pipeline.consumer_release(clc_pipe_throttle_consumer_state); + ++clc_pipe_throttle_consumer_state; + + // Query next clcID and update producer state + clc_pipe_producer_state = scheduler.advance_to_next_work(clc_pipeline, clc_pipe_producer_state); + } + + // Fetch next work tile + auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work( + work_tile_info, + clc_pipeline, + clc_pipe_consumer_state + ); + + // Only perform a new CLC query if we consumed a new CLC query result in + // `fetch_next_work`. An example of a case in which CLC `fetch_next_work` does + // not consume a new CLC query response is when processing stream-K units. + // The current stream-K scheduler uses single WorkTileInfo to track multiple + // (potentially-partial) tiles to be computed via stream-K. In this case, + // `fetch_next_work` simply performs in-place updates on the existing WorkTileInfo, + // rather than consuming a CLC query response. + requires_clc_query = increment_pipe; + if (increment_pipe) { + ++clc_pipe_consumer_state; + } + + work_tile_info = next_work_tile_info; + } while (work_tile_info.is_valid()); + clc_pipeline.producer_tail(clc_pipe_producer_state); + } + } + + else if (is_participant.mma) { + // Tmem allocation sequence + tmem_allocator.allocate(ArchTag::kTmemCapacityColumns, &shared_storage.tmem_base_ptr); + __syncwarp(); + tmem_allocation_result_barrier.arrive(); + uint32_t tmem_base_ptr = shared_storage.tmem_base_ptr; + collective_mainloop.set_tmem_offsets(tmem_storage, tmem_base_ptr); + + auto mma_inputs = collective_mainloop.mma_init( + tmem_storage, + shared_storage.tensors.mainloop); + + do { + auto k_tile_count = TileScheduler::get_work_k_tile_count(work_tile_info, problem_shape_MNKL, CtaShape_MNK{}); + + // Fetch next work tile + auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work( + work_tile_info, + clc_pipeline, + clc_pipe_consumer_state + ); + + if (increment_pipe) { + ++clc_pipe_consumer_state; + } + + // Accumulator stage slice + int acc_stage = accumulator_pipe_producer_state.index(); + + if (is_mma_leader_cta) { + mainloop_pipe_consumer_state = collective_mainloop.mma( + cute::make_tuple(mainloop_pipeline, accumulator_pipeline), + cute::make_tuple(mainloop_pipe_consumer_state, accumulator_pipe_producer_state), + collective_mainloop.slice_accumulator(tmem_storage, acc_stage), + mma_inputs, + cta_coord_mnkl, + k_tile_count + ); + accumulator_pipeline.producer_commit(accumulator_pipe_producer_state); + } + ++accumulator_pipe_producer_state; + + work_tile_info = next_work_tile_info; + cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info); + } while (work_tile_info.is_valid()); + + // Hint on an early release of global memory resources. + // The timing of calling this function only influences performance, + // not functional correctness. + cutlass::arch::launch_dependent_grids(); + // Release the right to allocate before deallocations so that the next CTA can rasterize + tmem_allocator.release_allocation_lock(); + + // Leader MMA waits for leader + peer epilogues to release accumulator stage + if (is_mma_leader_cta) { + accumulator_pipeline.producer_tail(accumulator_pipe_producer_state); + } + // Signal to peer MMA that entire tmem allocation can be deallocated + if constexpr (has_mma_peer_cta) { + // Leader does wait + arrive, follower does arrive + wait + tmem_deallocation_result_barrier.arrive(mma_peer_cta_rank, not is_mma_leader_cta); + tmem_deallocation_result_barrier.wait(dealloc_barrier_phase); + tmem_deallocation_result_barrier.arrive(mma_peer_cta_rank, is_mma_leader_cta); + } + + // Free entire tmem allocation + tmem_allocator.free(tmem_base_ptr, ArchTag::kTmemCapacityColumns); + } + + else if (is_participant.epi_load) { + // Ensure that the prefetched kernel does not touch + // unflushed global memory prior to this instruction + cutlass::arch::wait_on_dependent_grids(); + bool do_load_order_wait = true; + bool do_tail_load = false; + int current_wave = 0; + + do { + bool compute_epilogue = TileScheduler::compute_epilogue(work_tile_info, params.scheduler); + + // Get current work tile and fetch next work tile + auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work( + work_tile_info, + clc_pipeline, + clc_pipe_consumer_state + ); + work_tile_info = next_work_tile_info; + + if (increment_pipe) { + ++clc_pipe_consumer_state; + } + + if (compute_epilogue) { + if (do_load_order_wait) { + load_order_barrier.wait(); + do_load_order_wait = false; + } + + epi_load_pipe_producer_state = collective_epilogue.load( + epi_load_pipeline, + epi_load_pipe_producer_state, + problem_shape_MNKL, + CtaShape_MNK{}, + cta_coord_mnkl, + TileShape{}, + TiledMma{}, + shared_storage.tensors.epilogue + ); + + do_tail_load = true; + } + current_wave++; + + // Calculate the cta coordinates of the next work tile + cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info); + } while (work_tile_info.is_valid()); + + // Only perform a tail load if one of the work units processed performed + // an epilogue load. An example of a case in which a tail load should not be + // performed is in split-K if a cluster is only assigned non-final splits (for which + // the cluster does not compute the epilogue). + if (do_tail_load) { + collective_epilogue.load_tail( + epi_load_pipeline, epi_load_pipe_producer_state, + epi_store_pipeline, epi_store_pipe_producer_state); + } + } + + else if (is_participant.epilogue) { + // Wait for tmem allocate here + tmem_allocation_result_barrier.arrive_and_wait(); + uint32_t tmem_base_ptr = shared_storage.tmem_base_ptr; + collective_mainloop.set_tmem_offsets(tmem_storage, tmem_base_ptr); + + bool do_tail_store = false; + do { + // Fetch next work tile + auto [next_work_tile_info, increment_pipe] = scheduler.fetch_next_work( + work_tile_info, + clc_pipeline, + clc_pipe_consumer_state + ); + + if (increment_pipe) { + ++clc_pipe_consumer_state; + } + + // Accumulator stage slice + int acc_stage = accumulator_pipe_consumer_state.index(); + + auto accumulator = get<0>(collective_mainloop.slice_accumulator(tmem_storage, acc_stage)); + accumulator_pipe_consumer_state = scheduler.template fixup( + TiledMma{}, + work_tile_info, + accumulator, + accumulator_pipeline, + accumulator_pipe_consumer_state, + typename CollectiveEpilogue::CopyOpT2R{} + ); + + // + // Epilogue and write to gD + // + + // Take an input tensor of the form + // ((MMA_TILE_M, MMA_TILE_N), MMA_M, MMA_N, extra...) and turn it into: + // (((MMA_TILE_M, MMA_M), (MMA_TILE_N, MMA_N)), 1, 1, extra...) + auto zip_mma_modes = [](auto& tensor) { + auto L = tensor.layout(); + constexpr int R = cute::rank_v; + static_assert(R >= 3, "zip_mma_modes requires a tensor of rank >= 3"); + + auto core_shape = make_shape( + make_shape( + make_shape(get<0,0>(shape(L)), get<1>(shape(L))), // (MMA_TILE_M, MMA_M) + make_shape(get<0,1>(shape(L)), get<2>(shape(L)))), // (MMA_TILE_N, MMA_N) + Int<1>{}, + Int<1>{}); + + auto core_stride = make_stride( + make_stride( + make_stride(get<0,0>(stride(L)), get<1>(stride(L))), + make_stride(get<0,1>(stride(L)), get<2>(stride(L)))), + Int<0>{}, + Int<0>{}); + + auto new_shape = tuple_cat(core_shape, take<3, R>(shape(L))); + auto new_stride = tuple_cat(core_stride, take<3, R>(stride(L))); + + return make_tensor(tensor.data(), make_layout(new_shape, new_stride)); + }; + + if (scheduler.compute_epilogue(work_tile_info)) { + auto [load_state_next, store_state_next, acc_state_next] = collective_epilogue.store( + epi_load_pipeline, + epi_load_pipe_consumer_state, + epi_store_pipeline, + epi_store_pipe_producer_state, + accumulator_pipeline, + accumulator_pipe_consumer_state, + problem_shape_MNKL, + CtaShape_MNK{}, + cta_coord_mnkl, + TileShape{}, + TiledMma{}, + zip_mma_modes(accumulator), + shared_storage.tensors.epilogue + ); + epi_load_pipe_consumer_state = load_state_next; + epi_store_pipe_producer_state = store_state_next; + accumulator_pipe_consumer_state = acc_state_next; + do_tail_store = true; + } + work_tile_info = next_work_tile_info; + cta_coord_mnkl = scheduler.work_tile_to_cta_coord(work_tile_info); + + } while (work_tile_info.is_valid()); + + // Only perform a tail store if one of the work units processed performed + // an epilogue. An example of a case in which a tail load should not be + // performed is in split-K if a cluster is only assigned non-final splits (for which + // the cluster does not compute the epilogue). + if (do_tail_store) { + collective_epilogue.store_tail( + epi_load_pipeline, epi_load_pipe_consumer_state, + epi_store_pipeline, epi_store_pipe_producer_state, + CtaShape_MNK{}); + } + } + + else { + } + } +}; + +/////////////////////////////////////////////////////////////////////////////// + +} // namespace cutlass::gemm::kernel diff --git a/include/cutlass/gemm/kernel/tile_scheduler_params.h b/include/cutlass/gemm/kernel/tile_scheduler_params.h index 09fda77e13..50bc939ec3 100644 --- a/include/cutlass/gemm/kernel/tile_scheduler_params.h +++ b/include/cutlass/gemm/kernel/tile_scheduler_params.h @@ -73,7 +73,10 @@ get_max_cta_occupancy(int max_sm_per_gpc, GemmCoord cluster_shape, int sm_count) int const max_cta_occupancy_per_residual_gpc = num_gpc_residual - (num_gpc_residual % cluster_size); cta_per_device += max_cta_occupancy_per_residual_gpc; - cta_per_device = sm_count < cta_per_device ? sm_count : cta_per_device; + // Clamp to sm_count without breaking the whole-cluster rounding above. + if (sm_count < cta_per_device) { + cta_per_device = platform::max(cluster_size, (sm_count / cluster_size) * cluster_size); + } return cta_per_device; } diff --git a/include/cutlass/pipeline/pipeline.hpp b/include/cutlass/pipeline/pipeline.hpp index b9c04df49c..0f1a735113 100644 --- a/include/cutlass/pipeline/pipeline.hpp +++ b/include/cutlass/pipeline/pipeline.hpp @@ -33,6 +33,5 @@ //////////////////////////////////////////////////////////////////////////////////////////////////// #include "cutlass/pipeline/sm90_pipeline.hpp" -#include "cutlass/pipeline/sm100_pipeline.hpp" - +#include "cutlass/pipeline/sm100_pipeline.hpp" //////////////////////////////////////////////////////////////////////////////////////////////////// diff --git a/include/cutlass/pipeline/sm100_pipeline.hpp b/include/cutlass/pipeline/sm100_pipeline.hpp index 3892e42944..6c29318244 100644 --- a/include/cutlass/pipeline/sm100_pipeline.hpp +++ b/include/cutlass/pipeline/sm100_pipeline.hpp @@ -538,7 +538,7 @@ class PipelineTmaUmmaAsync { public: static constexpr uint32_t Stages = Stages_; using AtomThrShape_MNK = AtomThrShape_MNK_; -private: +protected: using Impl = PipelineTmaAsync; public: using FullBarrier = typename Impl::FullBarrier; @@ -734,6 +734,8 @@ class PipelineTmaUmmaAsync { uint16_t block_id_mask_ = 0; static constexpr bool is_2sm_mma = size(AtomThrShape_MNK{}) > 1; +private: + // Consumer signalling Producer of completion // Ensures all blocks in the Same Row and Column get notifed. CUTLASS_DEVICE diff --git a/include/cutlass/pipeline/sm90_pipeline.hpp b/include/cutlass/pipeline/sm90_pipeline.hpp index bcf6c743da..3dd68f90c9 100644 --- a/include/cutlass/pipeline/sm90_pipeline.hpp +++ b/include/cutlass/pipeline/sm90_pipeline.hpp @@ -202,7 +202,11 @@ struct PipelineState { CUTLASS_DEVICE void operator++() { - if constexpr (Stages > 0) { + if constexpr (Stages == 1) { + phase_ ^= 1; + ++count_; + } + else if constexpr (Stages > 0) { ++index_; ++count_; if (index_ == Stages) { @@ -243,6 +247,20 @@ struct PipelineState { return *this; } + template + CUTLASS_DEVICE + PipelineState& advance() { + if constexpr (Stages > 0 && NumIterations == Stages) { + // A full traversal returns to the same stage and flips the phase once. + phase_ ^= 1; + count_ += NumIterations; + return *this; + } + else { + return advance(NumIterations); + } + } + CUTLASS_DEVICE static PipelineState make_pipeline_state(PipelineState start_state, uint32_t num_iterations) { return start_state.advance(num_iterations); diff --git a/media/docs/cutlass_compiler/index.rst b/media/docs/cutlass_compiler/index.rst index 6e77e18dcd..f3344fc895 100644 --- a/media/docs/cutlass_compiler/index.rst +++ b/media/docs/cutlass_compiler/index.rst @@ -9,8 +9,8 @@ first-class MLIR types and operations. If you've worked with C++ CuTe, almost every concept here has a one-to-one analogue in CuTe IR. If you're new to CuTe, start with the -:doc:`Quickstart ` and then work through -the :doc:`tutorials`. +`Quickstart `__ and then work through +the `Tutorials `__. --------------------------------------------------------------------------- @@ -19,17 +19,17 @@ Overview This documentation is split into four sections: -- :doc:`Quickstart ` — what CuTe IR is, +- `Quickstart `__ — what CuTe IR is, how to read its syntax, and a minimal end-to-end example. Start here. -- :doc:`Tutorials ` — guided introduction to layouts and the +- `Tutorials `__ — guided introduction to layouts and the layout algebra. Read these first if unfamiliar with CuTe. -- :doc:`Cute dialect reference ` — exhaustive reference for +- `Cute dialect reference `__ — exhaustive reference for every op, type, and pass in the ``cute`` dialect. Covers operand types, assembly format, traits, and pass options. -- :doc:`Base dialect reference ` — the base facade and its +- `Base dialect reference `__ — the base facade and its target-attach / GPU-binary-emit passes (``attach-nvvm-target``, ``emit-gpu-binary``, ``base-prepare``, ``one-shot-convert-to-llvm``). diff --git a/media/docs/operators/api_reference/discovery.rst b/media/docs/operators/api_reference/discovery.rst index 7e463a5311..b0c5e4a7c7 100644 --- a/media/docs/operators/api_reference/discovery.rst +++ b/media/docs/operators/api_reference/discovery.rst @@ -11,6 +11,9 @@ Operators that support a given operation, operands, and target. It returns Operators can be additionally discovered or filtered by target compute capability, which is described by :class:`~cutlass.operators.TargetSm`. +For performance ranking and the ``heuristic=`` parameter, see +:doc:`heuristics`. + .. autofunction:: cutlass.operators.get_operators .. automodule:: cutlass.operators.manifest diff --git a/media/docs/operators/api_reference/heuristics.rst b/media/docs/operators/api_reference/heuristics.rst new file mode 100644 index 0000000000..f180958001 --- /dev/null +++ b/media/docs/operators/api_reference/heuristics.rst @@ -0,0 +1,108 @@ +.. _operators_api_reference_heuristics: + +Heuristics +========== + +Heuristics rank the Operators returned by +:func:`~cutlass.operators.get_operators` by estimated performance. Discovery +first filters candidates for correctness; a heuristic then orders and may +prune those candidates. A candidate's rank is its position in the returned +list. Ranking is an estimate, not a guarantee of the fastest Operator. + +For a worked GEMM example, see the +:doc:`heuristics tutorial `. + +Selecting a heuristic +--------------------- + +Pass a :class:`~cutlass.operators.Heuristic` instance as ``heuristic=`` to +:func:`~cutlass.operators.get_operators`. For example, given +:class:`~cutlass.operators.GemmArguments` named ``args``: + +.. code-block:: python + + import cutlass.operators as ops + from cutlass.operators.heuristics import NvMatmulHeuristics + + heuristic = NvMatmulHeuristics(gpu="B200") + # Equivalent construction through the registry: + heuristic = ops.get_heuristic("nvmatmul")(gpu="B200") + operators = ops.get_operators( + args, target_sm="100a", heuristic=heuristic, limit=5 + ) + +:func:`~cutlass.operators.get_heuristic` returns a class; instantiate it before +passing it to ``get_operators``. ``limit`` must be positive when provided. +It is forwarded to the heuristic as a hint, and ``get_operators`` truncates +the ranked result afterward. Pruning may leave fewer than ``limit`` results. +Omitting ``heuristic`` preserves discovery order. Ranking errors propagate +to the caller. + +Built-in nvMatmul heuristic +---------------------------- + +``NvMatmulHeuristics`` supports SM100 non-blockscaled dense GEMM. Its ``gpu`` +argument selects the modeled GPU SKU: ``"B200"`` (the default), +``"GB200_NVL"``, or ``"GB300_NVL"``. Other values raise ``ValueError``. +The model is selected explicitly, without detecting the current GPU. + +Only Operators designed for the modeled compute capability and matching a +recommended configuration are returned. Other generations and unmatched +Operators are excluded, including kernels designed for older generations +that can run on the target. ``target_sm`` filters discovery for compatibility; +it does not change the heuristic's GPU model. + +With non-empty candidates, unsupported arguments or layouts, missing +dependencies, and failure to match any recommendation raise errors. Empty +candidates return an empty list. + +Install the optional dependency with +``pip install 'nvidia-cutlass-operators[heuristics]'``. Registration does not +imply that this dependency is available: use +:func:`~cutlass.operators.heuristics.nvmatmul.is_available` to check package +compatibility. This check does not validate a particular GEMM or guarantee +that a candidate will match. + +.. autoclass:: cutlass.operators.heuristics.NvMatmulHeuristics + :class-doc-from: both + :members: rank + :show-inheritance: + +.. autofunction:: cutlass.operators.heuristics.nvmatmul.is_available + +.. py:data:: cutlass.operators.heuristics.nvmatmul.MIN_NVMMH_VERSION + :type: str + + Minimum supported version of the optional ``nvidia-matmul-heuristics`` + package. Availability also requires a compatible package API. + +Custom heuristics and registry +------------------------------ + +Subclass :class:`~cutlass.operators.Heuristic` and implement ``rank``. Return +an ordered subset of the input Operator objects; pruning is allowed. +``get_operators`` raises ``RuntimeError`` if ranking introduces an Operator +or duplicates one beyond its count in the input. ``limit`` is a hint to the +ranker; ``get_operators`` enforces the final result limit. + +Pass an instance directly, or register the class for lookup with +:func:`~cutlass.operators.get_heuristic`. Lookup returns the class itself and +raises ``KeyError`` for an unknown name. + +The symbols below are also exported by ``cutlass.operators.heuristics``. +The nvMatmul ``_mapping`` and ``_provider`` modules are implementation details; +their query, configuration, and matching helpers are not public APIs. + +.. autoclass:: cutlass.operators.Heuristic + :members: rank + +.. autofunction:: cutlass.operators.register_heuristic + +.. autofunction:: cutlass.operators.get_heuristic + +.. py:data:: cutlass.operators.available_heuristics + :type: dict[str, type[Heuristic]] + + Maps registered names to heuristic classes. The built-in ``"nvmatmul"`` + class registers even when its optional dependency is absent. Registering an + existing name replaces its previous class. diff --git a/media/docs/operators/api_reference/index.rst b/media/docs/operators/api_reference/index.rst index fc34722f6b..d7ce69c2db 100644 --- a/media/docs/operators/api_reference/index.rst +++ b/media/docs/operators/api_reference/index.rst @@ -16,5 +16,6 @@ The API reference is grouped into the following pages: operator arguments discovery + heuristics metadata misc diff --git a/media/docs/operators/overview.rst b/media/docs/operators/overview.rst index 3c23ce575d..126d3cb1c5 100644 --- a/media/docs/operators/overview.rst +++ b/media/docs/operators/overview.rst @@ -126,14 +126,17 @@ CUTLASS Operator API will support a wide range of functionality, configurations, - Kernel coverage - - Dense GEMMs (F32, F16, BF16, INT8) for Blackwell, Hopper, Ampere + - Dense GEMMs for Rubin, Blackwell, Hopper, Ampere + + - FP8 for Rubin, with more in future. Various combinations of F32, F16, BF16, FP8, INT8 for Blackwell and older. - Preferred and fallback cluster shapes - Static and dynamic scheduling - - Block-scaled GEMMs (NVFP4, MXFP4, MXFP8, mixed input precision) for Blackwell + + - Block-scaled GEMMs (NVFP4, MXFP4, MXFP8, mixed input precision) for Rubin and Blackwell - Grouped GEMM (Contiguous offset) for Blackwell - Low-latency TGV GEMM for Blackwell -- Custom epilogue fusions +- Custom epilogue fusions for Blackwell - Activations, elementwise ops, auxiliary tensor load/store - Row/column broadcasts @@ -143,13 +146,14 @@ CUTLASS Operator API will support a wide range of functionality, configurations, - CUDA Graph support - Native support for PyTorch and other DLPack tensors - Bring-your-own-kernel +- Operator ranking via nvMatmulHeuristics (``heuristic=NvMatmulHeuristics(gpu=...)``) + for dense SM100 GEMM **Upcoming support:** - Additional GEMM kernel coverage: Sparsity, performance optimizations, grouped GEMM variants, and more - Ahead-of-time compilation - JAX Graph support -- nvMatmulHeuristics support Community & Feedback ==================== diff --git a/media/docs/pythonDSL/cute_dsl_general/dsl_control_flow.rst b/media/docs/pythonDSL/cute_dsl_general/dsl_control_flow.rst index ee7d9200e4..d4e15e5df0 100644 --- a/media/docs/pythonDSL/cute_dsl_general/dsl_control_flow.rst +++ b/media/docs/pythonDSL/cute_dsl_general/dsl_control_flow.rst @@ -124,13 +124,15 @@ Compiler will automatically generate the prefetch loop with `prefetch_stages` it .. note:: - The compiler splits the loop at a TMA bulk copy: the loop body must - contain at least one ``cute.copy(tma_atom, ..., tma_bar_ptr=...)`` - paired with an mbarrier wait on the same barrier. Loops built from - plain ``cute.copy``, ``cp.async``, or arithmetic currently have no - such split point: the compiler emits a warning - ("software pipelining ('prefetch_stages') skipped") and compiles the - loop without pipelining. + The compiler recognizes two split points. A TMA bulk copy: at least one + ``cute.copy(tma_atom, ..., tma_bar_ptr=...)`` paired with an mbarrier + wait on the same barrier. Or a cp.async group: copies with a cp.async + atom followed by ``cute.arch.cp_async_commit_group()`` and + ``cute.arch.cp_async_wait_group(0)``, in a canonical loop (start 0, + step 1) whose circular-buffer slot is indexed as ``i % stages``. + Loops with neither split point — or outside these shapes — get a + warning ("software pipelining ('prefetch_stages') skipped") and are + compiled without pipelining. This feature is experimental and only supported on sm90 and above. diff --git a/media/docs/pythonDSL/cute_dsl_general/dsl_jit_compilation_options.rst b/media/docs/pythonDSL/cute_dsl_general/dsl_jit_compilation_options.rst index 400f56c538..fd0622d0da 100644 --- a/media/docs/pythonDSL/cute_dsl_general/dsl_jit_compilation_options.rst +++ b/media/docs/pythonDSL/cute_dsl_general/dsl_jit_compilation_options.rst @@ -24,7 +24,7 @@ The |DSL| provides multiple ways to specify compilation options - either by spec ``cute.compile`` Compilation Options as strings ----------------------------------------------- -You can provide additional compilation options as a string when calling ``cute.compile``. The |DSL| uses ``argparse`` to parse these options and will raise an error if any invalid options are specified. +You can provide additional compilation options as a string when calling ``cute.compile``. The |DSL| uses ``argparse`` to parse the documented driver options below and will raise an error if any invalid options are specified. .. list-table:: :header-rows: 1 @@ -92,7 +92,6 @@ You can use the following code to specify compilation options: jit_executor_with_nvdisasm_options = cute.compile(add, 1, 2, options="--keep-sass --nvdisasm-options '-c'") jit_executor_with_ptxas_options = cute.compile(add, 1, 2, options="--ptxas-options '--opt-level=2'") - ``cute.compile`` Compilation Options as separate Python types ------------------------------------------------------------- @@ -101,7 +100,16 @@ Compilation options can be programmatically composed using tuple and passed to ` .. code-block:: python - from cutlass.cute import OptLevel, EnableAssertions, GenerateLineInfo, KeepCUBIN, KeepPTX, KeepSASS, NvdisasmOptions + from cutlass.cute import ( + OptLevel, + EnableAssertions, + GenerateLineInfo, + KeepCUBIN, + KeepPTX, + KeepSASS, + NvdisasmOptions, + PtxasOptions, + ) my_debugging_options = (OptLevel(1), EnableAssertions, GenerateLineInfo, KeepCUBIN, KeepPTX) compiled_kernel_1 = cute.compile[my_debugging_options](my_kernel_1, ...) diff --git a/media/docs/pythonDSL/guides/ahead_of_time_compilation.rst b/media/docs/pythonDSL/guides/ahead_of_time_compilation.rst index ce7939978a..24f5a63d10 100644 --- a/media/docs/pythonDSL/guides/ahead_of_time_compilation.rst +++ b/media/docs/pythonDSL/guides/ahead_of_time_compilation.rst @@ -141,6 +141,7 @@ Dynamically load pre-compiled object files or shared libraries at runtime. By in #include "CuteDSLRuntime.h" #include + #include void run_print_tensor() { // Load module from shared library @@ -179,9 +180,14 @@ Dynamically load pre-compiled object files or shared libraries at runtime. By in cudaStreamCreate(&stream); // Call the function; the runtime function accepts packed arguments, refer to the wrapper in the header file - int ret; + // The trailing packed argument receives the CUDA error code of the kernel launch + int32_t ret = 0; void* args[] = {&tensor_a, &stream, &ret}; err = CuteDSLRT_Function_Run(func, args, 3); + if (ret != cudaSuccess) { + fprintf(stderr, "kernel launch failed: %s\n", + cudaGetErrorName(static_cast(ret))); + } check_error(err); cudaStreamSynchronize(stream); @@ -192,7 +198,7 @@ Dynamically load pre-compiled object files or shared libraries at runtime. By in The ``CuteDSLRuntime.h`` header file can be found in ``/include``. It includes: -* The ``CuteDSLRT_Error_t`` type: Indicates error status. +* The ``CuteDSLRT_Error_t`` type: Indicates the status of the runtime API call itself, not the CUDA error code of the kernel launch. See :ref:`dsl_aot_error_handling`. * The ``CuteDSLRT_Module_Load`` function: Loads the module. * The ``CuteDSLRT_Module_Get_Function`` function: Gets a function from the loaded module. The runtime API will load the CUDA module for kernel execution. * The ``CuteDSLRT_Function_Run`` function: Runs the function. @@ -200,6 +206,68 @@ The ``CuteDSLRuntime.h`` header file can be found in ``/incl The compilation of the C++ executable requires the ``libcute_dsl_runtime.so`` library which is involved in ``/lib``, along with the CUDA driver and runtime libraries, to function properly. +.. _dsl_aot_error_handling: + +Return Values and Error Handling +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The wrapper function in the generated header returns an ``int32_t``: + +.. code-block:: cpp + + static inline int32_t cute_dsl_print_tensor_wrapper( + print_tensor_Kernel_Module_t *module, + print_tensor_Tensor_a_t *a, + cudaStream_t stream); + +The returned value is a CUDA runtime ``cudaError_t`` code: + +* ``0`` (``cudaSuccess``) means every kernel launch in the exported function was + submitted successfully. +* Any other value is the code returned by ``cudaLaunchKernelExC`` for the first + launch that failed. The exported function returns at that point, so kernels + launched later in the same ``@cute.jit`` function do not run. + +Because the value is an ordinary ``cudaError_t``, the CUDA runtime helpers +``cudaGetErrorName`` and ``cudaGetErrorString`` translate it into a +human-readable message. The generated header also defines a +``CUTE_DSL_CUDA_ERROR_CHECK`` macro, defined in terms of those two helpers, +that reports the code this way: + +.. code-block:: cpp + + #include "print_tensor_example.h" + + int32_t ret = cute_dsl_print_tensor_wrapper(&module, &tensor_a, stream); + CUTE_DSL_CUDA_ERROR_CHECK(ret); + + // Or inspect the code directly + if (ret != cudaSuccess) { + cudaError_t err = static_cast(ret); + fprintf(stderr, "kernel launch failed: %s: %s\n", + cudaGetErrorName(err), cudaGetErrorString(err)); + } + +Note that: + +* **The return value only covers launch submission.** Kernel execution is + asynchronous, so faults such as ``cudaErrorIllegalAddress`` do not appear in + it. Check ``cudaStreamSynchronize`` or ``cudaDeviceSynchronize`` separately + for those. +* **Dynamic loading reports two independent statuses.** ``CuteDSLRT_Error_t`` + describes the runtime API call itself, such as module loading, symbol lookup + and invocation, and is decoded with ``CuteDSLRT_GetErrorName`` and + ``CuteDSLRT_GetErrorString``. It reports every kernel failure as the single + value ``CuteDSLRT_Error_CudaError`` and does not carry the underlying code; + that code is written to the trailing packed argument instead, as shown in the + dynamic loading example above. +* **Loading in Python raises instead of returning.** A non-zero code is raised + as a ``DSLCudaRuntimeError`` carrying the ``cudaError_t`` name. +* **The Apache TVM FFI ABI uses a different contract.** Its exported functions + return ``0`` on success and ``-1`` on failure, and the message is retrieved + through the TVM FFI error object rather than from the return value. See + :doc:`tvm_ffi_compilation`. + .. _dsl_aot_host_cross_compilation: Host Cross-Compilation for AArch64 @@ -288,11 +356,81 @@ Linking succeeds against the stub on the build host. Deploy ``kernel.so`` to the Limitations ~~~~~~~~~~~ -* **AArch64 only.** Only ``aarch64-unknown-linux-gnu`` is currently supported. Other architectures fail with a "target not registered" error during code generation. +* **AArch64 only.** Only ``aarch64-unknown-linux-gnu`` and ``aarch64-unknown-nto-qnx8.0.0`` (see below) are currently supported. Other architectures fail with a "target not registered" error during code generation. * **Not compatible with TVM FFI.** Combining ``--enable-tvm-ffi`` with ``--host-target`` raises an error; drop ``--enable-tvm-ffi`` and use the plain AOT export path described here. * **CUDA runtime version must match.** The exported object depends on the CUDA runtime; the target's CUDA runtime/toolkit version must match the one |DSL| was built against (see :ref:`dsl_aot_object_compat`). * **Linking is your responsibility.** You must supply your own cross toolchain and a target sysroot with the CUDA headers and libraries; the stub only resolves the |DSL| runtime symbols. +Host Cross-Compilation for QNX 8.0 +---------------------------------- + +The same AOT export targets QNX 8.0 on AArch64. Only the **static-linking** +integration is supported: the exported object is linked into your final QNX +shared library or executable at build time. Dynamic loading +(``CuteDSLRT_Module_*``, ``cute.runtime.load_module``) +is not available on QNX, because the module loader is built on LLVM ORC JIT, +which is not cross-built for that platform. Those entry points still exist and +return ``CuteDSLRT_Error_UnsupportedOnPlatform``. + +Select the target with the ``qnx8-aarch64`` preset:: + + compiled = cute.compile( + my_function, *args, + options="--gpu-arch sm_110a --host-target qnx8-aarch64") + + compiled.export_to_c(file_path="./artifacts", file_name="kernel", + function_prefix="kernel") + +The preset maps to the triple ``aarch64-unknown-nto-qnx8.0.0``. LLVM has no QNX +target, so that triple resolves to an unknown OS and generic AArch64 ELF +codegen: the emitted object is identical to the one ``linux-aarch64`` produces. +The distinct triple exists so that the runtime-library lookup below can tell the +two targets apart. Because the object is plain AArch64 ELF and QNX uses the same +AAPCS64 ABI, the QNX linker consumes it directly. + +Cross-link with the QNX 8.0 toolchain:: + + q++ -Vgcc_ntoaarch64le -shared -o kernel.so kernel.o \ + $(python -m cutlass.cute.export.aot_config --ldflags --target aarch64-unknown-nto-qnx8.0.0) \ + $(python -m cutlass.cute.export.aot_config --libs --target aarch64-unknown-nto-qnx8.0.0) + +Obtaining the Target Runtime +~~~~~~~~~~~~~~~~~~~~~~~~~~~~~ + +The wheel by default ships only the **link-time stub** for QNX, under +``lib/stubs/aarch64-unknown-nto-qnx8.0.0/``. It resolves the |DSL| runtime +symbols on the build host and is never executed. This mirrors how the CUDA +toolkit ships ``lib/stubs/libcuda.so``. + +The real ``libcute_dsl_runtime.so`` for QNX is obtained separately via the extras +``nvidia-cutlass-dsl[qnx]`` or ``nvidia-cutlass-dsl[qnx-cu13]`` on x86_64 Linux, and +is installed at ``lib/aarch64-unknown-nto-qnx8.0.0/``. It is only +needed on the target: deploy it into the QNX root filesystem as part of your +image build, where the dynamic loader binds against it at run time. + +Version Compatibility +~~~~~~~~~~~~~~~~~~~~~~ + +Because the runtime and the wheel are obtained separately, they must be matched +explicitly: + +* The QNX runtime artifact must come from the **same |DSL| build** as the wheel + that produced the object. The object embeds a version that is checked when a + module is loaded, and the two are only guaranteed consistent when they are + built together. +* The QNX CUDA toolkit on the target must match the CUDA version |DSL| was built + against, exactly as in the AArch64 Linux case + (see :ref:`dsl_aot_object_compat`). + +QNX Limitations +~~~~~~~~~~~~~~~~ + +* **Static linking only.** Dynamic loading is unavailable; see above. +* **AArch64 only.** QNX on other architectures is not supported. +* **Restricted CUDA surface.** A safety or otherwise restricted QNX CUDA build + may not provide every CUDA entry point the |DSL| runtime references, which + surfaces as an unresolved symbol when linking the runtime for QNX. + Supported Argument Types ------------------------ diff --git a/media/docs/pythonDSL/guides/debugging.rst b/media/docs/pythonDSL/guides/debugging.rst index f20b02b774..9b84a6a89d 100644 --- a/media/docs/pythonDSL/guides/debugging.rst +++ b/media/docs/pythonDSL/guides/debugging.rst @@ -139,6 +139,11 @@ category. - ``ptxas`` resource diagnostics surfaced through the remark stream, including register spills and local-memory usage. - Info (remark) + * - ``loop`` + - ``remarks{loop}`` + - Loop-optimization remarks, such as software pipelining and loop + unrolling applied by the compiler. + - Info (remark) For example, enable NVVM primitive diagnostics with: diff --git a/media/docs/pythonDSL/limitations.rst b/media/docs/pythonDSL/limitations.rst index 14b112af9b..99dbb92745 100644 --- a/media/docs/pythonDSL/limitations.rst +++ b/media/docs/pythonDSL/limitations.rst @@ -117,16 +117,24 @@ Programming Model - An ``int`` outside the type's range **wraps**, dropping the high bits. For example, ``Int32(1 << 34)`` yields ``0`` and ``Int32(2**31)`` yields ``-2147483648``. - - A ``float`` whose magnitude is too large **overflows to** ``±inf``, and one - too small **underflows to** ``0``. For example, ``Float32(1e40)`` yields - ``inf`` and ``Float32(1e-50)`` yields ``0.0``. + - A ``float`` narrowed to an **integer** type truncates toward zero, and a + value the type cannot hold is then narrowed by the **host CPU's own + float-to-integer conversion** — the same cast NumPy performs for that + value on that machine. That result is architecture-defined: + ``Int32(1e40)`` yields ``-2147483648`` on x86-64 and ``2147483647`` on + AArch64. Treat it as unspecified and clamp before converting. + - A ``float`` narrowed to a **float** type whose magnitude is too large + **overflows to** ``±inf``, and one too small **underflows to** ``0``. For + example, ``Float32(1e40)`` yields ``inf`` and ``Float32(1e-50)`` yields + ``0.0``. Note the contrast: ``(1 << 34) + Int32(3)`` promotes to ``Int64`` and keeps the value, whereas the explicit ``Int32(1 << 34)`` wraps to ``0``. To surface this loss, the DSL emits a compiler **warning** for these catastrophic construction cases (``TYPE_INT_LITERAL_OUT_OF_RANGE``, - ``TYPE_FLOAT_LITERAL_OVERFLOW``, ``TYPE_FLOAT_LITERAL_UNDERFLOW``), pointing at + ``TYPE_FLOAT_TO_INT_OUT_OF_RANGE``, ``TYPE_FLOAT_LITERAL_OVERFLOW``, + ``TYPE_FLOAT_LITERAL_UNDERFLOW``), pointing at the exact source location and suggesting a wider type. Note that ordinary precision/rounding loss is *not* flagged, since it is inherent to every float literal and would be too noisy. For example, ``Float32(0.1)`` is @@ -142,10 +150,12 @@ Programming Model b = cutlass.Float32(1e40) # warns: overflows to inf c = cutlass.Float32(1e-50) # warns: underflows to 0.0 d = cutlass.Float32(0.1) # no warning: ordinary rounding + e = cutlass.Int32(1e40) # warns: host-defined narrowing Prefer a wider type (e.g. ``Int64`` / ``Float64``) when the value does not fit the constructor's type, or mask an ``int`` to the type width to make an - intentional wrap explicit. + intentional wrap explicit. Masking has no float spelling, so clamp a float + to ``[min, max]`` instead. **Python Function** The DSL currently has **limited support for return values** from Python functions. diff --git a/operators/README.md b/operators/README.md index 80ceaa8e81..1aabb1feae 100644 --- a/operators/README.md +++ b/operators/README.md @@ -120,14 +120,15 @@ CUTLASS Operator API will support a wide range of functionality, configurations, - Kernel coverage - - Dense GEMMs (F32, F16, BF16, INT8) for Blackwell, Hopper, Ampere + - Dense GEMMs for Rubin, Blackwell, Hopper, Ampere + - FP8 for Rubin, with more in future. Various combinations of F32, F16, BF16, FP8, INT8 for Blackwell and older. - Preferred and fallback cluster shapes - Static and dynamic scheduling - - Block-scaled GEMMs (NVFP4, MXFP4, MXFP8, mixed input precision) for Blackwell + - Block-scaled GEMMs (NVFP4, MXFP4, MXFP8, mixed input precision) for Rubin and Blackwell - Grouped GEMM (Contiguous offset) for Blackwell - Low-latency TGV GEMM for Blackwell -- Custom epilogue fusions +- Custom epilogue fusions for Blackwell - Activations, elementwise ops, auxiliary tensor load/store - Row/column broadcasts diff --git a/operators/cutlass/operators/heuristics/base.py b/operators/cutlass/operators/heuristics/base.py index 44dc3af168..05b8f9d1ff 100644 --- a/operators/cutlass/operators/heuristics/base.py +++ b/operators/cutlass/operators/heuristics/base.py @@ -69,8 +69,9 @@ def rank( Args: args (RuntimeArguments | None): Runtime arguments describing the - problem being ranked for (e.g. a :class:`GemmArguments`). May be - ``None`` if the caller did not provide arguments. + problem being ranked for (e.g. + :class:`~cutlass.operators.GemmArguments`). May be ``None`` if + the caller did not provide arguments. operators (list[Operator]): The already-filtered candidate Operators to order. Every element is known to support ``args``. target_sm (TargetSm | str | None): Optional compute capability the diff --git a/operators/cutlass/operators/heuristics/nvmatmul/ranker.py b/operators/cutlass/operators/heuristics/nvmatmul/ranker.py index 171146ffff..ba3b19fcb0 100644 --- a/operators/cutlass/operators/heuristics/nvmatmul/ranker.py +++ b/operators/cutlass/operators/heuristics/nvmatmul/ranker.py @@ -99,11 +99,13 @@ def is_available() -> bool: Mirrors the availability signalling of ``cutlass.operators.available_providers``: the heuristic is always registered, but this reports whether it can actually - run in this environment (importable, ``>={MIN_NVMMH_VERSION}``, and the - 0.1.0.27 constructor API shape). + run in this environment (importable, at least + :data:`~cutlass.operators.heuristics.nvmatmul.MIN_NVMMH_VERSION`, and + compatible with the 0.1.0.27 constructor API shape). Returns: - bool: ``True`` if a call to :meth:`NvMatmulHeuristics.rank` can use + bool: ``True`` if a call to + :meth:`~cutlass.operators.heuristics.NvMatmulHeuristics.rank` can use nvMMH, ``False`` if it would raise :class:`ImportError`. """ try: @@ -131,7 +133,8 @@ class NvMatmulHeuristics(Heuristic): Queries nvMMH for recommended configs, ranks Operators that match the recommended configs, and prunes away Operators that do not match any recommended config or that - do not exactly match the GPU modeled in `self.gpu`. + were designed for a different compute capability than the GPU selected by + ``gpu``. Missing optional dependencies raise :class:`ImportError`. Non-GEMM ``args`` raise :class:`TypeError`. @@ -174,7 +177,7 @@ def rank( ) -> list[Operator]: """Order ``operators`` best-first for ``args`` using nvMatmulHeuristics. - See :meth:`cutlass.operators.heuristics.base.Heuristic.rank`. + See :meth:`cutlass.operators.Heuristic.rank`. Args: args (RuntimeArguments | None): The problem to rank for. Must be @@ -182,7 +185,7 @@ def rank( non-empty. operators (list[Operator]): Filtered candidate Operators. target_sm (TargetSm | str | None): Accepted for - :meth:`Heuristic.rank` compatibility; unused. + :meth:`~cutlass.operators.Heuristic.rank` compatibility; unused. limit (int | None): When set, the caller will keep only the first ``limit`` Operators from the ranked result. diff --git a/operators/cutlass/operators/providers/cutedsl/gemm/implementations/sm100_tgv_gemm_impl.py b/operators/cutlass/operators/providers/cutedsl/gemm/implementations/sm100_tgv_gemm_impl.py index 6040d6693b..a82edb47a2 100644 --- a/operators/cutlass/operators/providers/cutedsl/gemm/implementations/sm100_tgv_gemm_impl.py +++ b/operators/cutlass/operators/providers/cutedsl/gemm/implementations/sm100_tgv_gemm_impl.py @@ -359,7 +359,7 @@ class SharedStorage: # Barrier between MMA and epilog: sync TMEM allocation/deallocation status tmem_allocation_result_barrier: cutlass.Int64 # Base pointer for TMEM allocation, MMA will write the allocated address here - tmem_base_ptr: cutlass.Int32 + tmem_base: cutlass.Int32 smem = cutlass.memory.SmemAllocator() storage = smem.allocate(SharedStorage) @@ -370,7 +370,7 @@ class SharedStorage: tma_epilog_full_bar = storage.tma_epilog_full_barrier.ptr mma_epilog_full_bar = storage.mma_epilog_full_barrier.ptr tmem_alloc_result_bar = storage.tmem_allocation_result_barrier.ptr - tmem_base_smem_ptr = storage.tmem_base_ptr.ptr + tmem_base_smem_ptr = storage.tmem_base.ptr # ============================================================ # Barrier initialization — ALL threads reach here, elect_one inits diff --git a/python/CuTeDSL/_mlir_helpers/arith.py b/python/CuTeDSL/_mlir_helpers/arith.py index 18c9efdaf4..a6afabc945 100644 --- a/python/CuTeDSL/_mlir_helpers/arith.py +++ b/python/CuTeDSL/_mlir_helpers/arith.py @@ -14,9 +14,8 @@ """ import array -from typing import Any, Callable, Optional, Union - -import numpy as np +import sys +from typing import Any, Callable, Optional, Union, TYPE_CHECKING from ..base_dsl.common import * from ..base_dsl.common import DSLRuntimeError, DSLNotImplemented @@ -27,6 +26,24 @@ from .lru_cache_ir import lru_cache_ir +if TYPE_CHECKING: + import numpy as np + from ..base_dsl.typing import Numeric + + +def _as_numpy_ndarray(value: Any) -> Any: + """Return ``value`` if it is a NumPy array, else None. + + NumPy is an optional dependency. A value cannot be an ``ndarray`` unless + NumPy is already imported, so consulting ``sys.modules`` is an exact test + that never imports NumPy on its behalf. + """ + numpy = sys.modules.get("numpy") + if numpy is None or not isinstance(value, numpy.ndarray): + return None + return value + + # ============================================================================= # Arith Dialect Helper functions # ============================================================================= @@ -342,9 +359,8 @@ def _cast( ) -@lru_cache_ir() def const( - value: Union[int, float, bool, np.ndarray], + value: Union[int, float, bool, "np.ndarray", "Numeric", "ArithValue"], ty: Optional[Union[ir.Type, "NumericMeta"]] = None, # type: ignore[name-defined] *, signed: Union[bool, None] = None, @@ -354,21 +370,33 @@ def const( """ Generates dynamic expression for constant values. """ - from ..base_dsl.typing import Numeric, NumericMeta - from ..base_dsl.dsl import is_dynamic_expression - from ..base_dsl.utils.numpy import _numpy_type_to_mlir_type - - # D1: ``_WatchedM`` is the PyIR meta-promotion wrapper. When fed into - # ``arith.const`` we ask the wrapper to bake its leaf (records the - # constant under the slot so a later mutation can rewrite it via - # pyir.load %ref). Lazy-import to keep PyIR-disabled builds clean. try: from ..base_dsl.pyir_runtime import _WatchedM - except ImportError: + except ModuleNotFoundError as e: + # The relative import always resolves to .base_dsl.pyir_runtime, + # so match the full suffix: a bare-basename match would also swallow a + # failing transitive dep that merely SHARES the module name. + if e.name is None or not e.name.endswith(".base_dsl.pyir_runtime"): + raise _WatchedM = None # type: ignore[misc,assignment] if _WatchedM is not None and isinstance(value, _WatchedM): return value.ir_value() + return _const_cached(value, ty, signed=signed, loc=loc, ip=ip) + + +@lru_cache_ir() +def _const_cached( + value: Union[int, float, bool, "np.ndarray", "Numeric", "ArithValue"], + ty: Optional[Union[ir.Type, "NumericMeta"]] = None, # type: ignore[name-defined] + *, + signed: Union[bool, None] = None, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, +) -> ir.Value: + from ..base_dsl.typing import Numeric, NumericMeta + from ..base_dsl.dsl import is_dynamic_expression + if isinstance(value, Numeric): value = value.value @@ -387,9 +415,11 @@ def const( ty = T.bool() elif isinstance(value, int): ty = T.i32() - elif isinstance(value, np.ndarray): - ty = T.vector(*value.shape, _numpy_type_to_mlir_type(value.dtype)) - value = array.array(value.dtype.kind, value.flatten().tolist()) # type: ignore[assignment] + elif (ndarray := _as_numpy_ndarray(value)) is not None: + from ..base_dsl.utils.numpy import _numpy_type_to_mlir_type + + ty = T.vector(*ndarray.shape, _numpy_type_to_mlir_type(ndarray.dtype)) + value = array.array(ndarray.dtype.kind, ndarray.flatten().tolist()) # type: ignore[assignment] else: raise DSLNotImplemented(f"{type(value)} is not supported") elif isinstance(ty, NumericMeta): diff --git a/python/CuTeDSL/_mlir_helpers/dominance.py b/python/CuTeDSL/_mlir_helpers/dominance.py new file mode 100644 index 0000000000..9b07c82071 --- /dev/null +++ b/python/CuTeDSL/_mlir_helpers/dominance.py @@ -0,0 +1,75 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: LicenseRef-NvidiaProprietary +# +# Use of this software is governed by the terms and conditions of the +# NVIDIA End User License Agreement (EULA), available at: +# https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/license.html +# +# Any use, reproduction, disclosure, or distribution of this software +# and related documentation outside the scope permitted by the EULA +# is strictly prohibited. + +"""IR-position validity checks for cached SSA values. + +An emitted op is *positional*: its results are only usable where its block +dominates. A Python wrapper that caches an ``ir.Value`` (or anything built +from one) and outlives the region it was built in can therefore serve an +operand that is invalid at a later use site. These helpers let such caches +verify a stored value before serving it, mirroring how ``lru_cache_ir`` keys +on the insertion point instead of assuming position-independence. +""" + +from typing import Any, Iterator + +from .._mlir import ir + +# The dominance query is a C++ binding that currently ships in the pyir +# dialect module. It is a generic MLIR facility (DominanceInfo), not a +# frontend feature; not every ``_mlir`` bundle carries the module, so its +# absence degrades to "assume reachable". +try: + from .._mlir.dialects import pyir as _pyir_dialect +except ImportError: + _pyir_dialect = None + + +def ssa_leaves(obj: Any) -> Iterator["ir.Value"]: + """Yield the SSA values backing *obj*. + + Covers the shapes derived caches take in practice: a raw ``ir.Value``, a + wrapper exposing ``.value``, and nested tuples/lists whose leaves may be + plain Python numbers carrying no SSA at all. + """ + if isinstance(obj, ir.Value): + yield obj + elif isinstance(obj, (tuple, list)): + for item in obj: + yield from ssa_leaves(item) + else: + inner = getattr(obj, "value", None) + if isinstance(inner, ir.Value): + yield inner + + +def value_reaches_current_ip(obj: Any) -> bool: + """Whether every SSA leaf of *obj* is usable at the current insertion + point. True when undecidable, so callers drop a cached value only once + the escape is positively proven. + + The insertion point's ref operation is part of the query: insertion may + sit *before* an existing op (e.g. ``InsertionPoint.at_block_begin``), and + a same-block value defined after that position must not count as + reachable.""" + if _pyir_dialect is None: + return True + try: + ip = ir.InsertionPoint.current + ref_op = ip.ref_operation + if ref_op is not None: + ref_op = getattr(ref_op, "operation", ref_op) + return all( + _pyir_dialect.value_dominates_ip(leaf, ip.block, ref_op) + for leaf in ssa_leaves(obj) + ) + except (RuntimeError, ValueError): + return True diff --git a/python/CuTeDSL/_mlir_helpers/op.py b/python/CuTeDSL/_mlir_helpers/op.py index 951c77e014..c060bbcfb7 100644 --- a/python/CuTeDSL/_mlir_helpers/op.py +++ b/python/CuTeDSL/_mlir_helpers/op.py @@ -21,8 +21,10 @@ import os import types import warnings +from contextlib import contextmanager +from contextvars import ContextVar from functools import wraps, lru_cache -from typing import Any, Callable, TYPE_CHECKING +from typing import Any, Callable, Iterator, TYPE_CHECKING from .._mlir import ir from ..base_dsl.common import ( @@ -45,11 +47,10 @@ _pyir_auto_load_arg, _pyir_lookup_slot_from_value, _pyir_value_tracked_by_accessible_ref, - _describe_value_origin, ) except ImportError: - def _pyir_auto_load_arg(arg: Any) -> Any: + def _pyir_auto_load_arg(arg: Any, *, row_authoritative: bool = False) -> Any: # noqa: ARG001 return arg def _pyir_lookup_slot_from_value(value: Any) -> "MutableValue | None": # noqa: ARG001 @@ -58,9 +59,6 @@ def _pyir_lookup_slot_from_value(value: Any) -> "MutableValue | None": # noqa: def _pyir_value_tracked_by_accessible_ref(value: Any) -> bool: # noqa: ARG001 return False - def _describe_value_origin(raw: Any) -> str: # noqa: ARG001 - return "" - # The DSL package root is empty by default. _DSL_PACKAGE_ROOT: str | None = "" @@ -69,6 +67,13 @@ def _describe_value_origin(raw: Any) -> str: # noqa: ARG001 # Whether location tracking is enabled. _ENABLE_FRAME_FILTERING: bool = False +# Generic stack of loc transforms used by dialect-specific compilation +# contexts. Keep this file unaware of any particular dialect or debug-info +# schema; callers decide what a transformed loc means. +_LOC_TRANSFORMS: ContextVar[tuple[Callable[[Any], Any], ...]] = ContextVar( + "_LOC_TRANSFORMS", default=() +) + # When True, dsl_user_op attributes ops to the closest frame (including DSL # library code) instead of skipping up to the user's call site. Enabled when # debugging mode is ON so library developers can see where inside the DSL an op @@ -135,6 +140,52 @@ def _set_enable_frame_filtering(enable: bool) -> None: _ENABLE_FRAME_FILTERING = enable +@contextmanager +def loc_transform(transform: Callable[[Any], Any]) -> Iterator[None]: + """Temporarily rewrite locs produced by ``@dsl_user_op``. + + The decorator still owns the generic work: find the Python user frame and + build the usual MLIR source loc. This hook lets a frontend or dialect wrap + that loc while it is building a scoped construct. + + Example: + def add_scope(loc): + if loc is None: + return None + return ir.Location.name("frontend.scope", childLoc=loc) + + with loc_transform(add_scope): + dsl_add(a, b) # receives loc=frontend.scope("file.py":line:col) + + Transforms are scoped and stackable. The innermost transform runs first, so + nested frontend scopes behave like nested Python context managers. + + ``transform`` receives each generated or caller-provided location, which may + be ``None``, and returns its replacement. Exceptions raised by ``transform`` + propagate to the caller. The active transform stack is restored when this + context exits, including exceptional exits. + + :param transform: Callable used to rewrite locations within this scope. + :type transform: Callable[[Any], Any] + :raises TypeError: If ``transform`` is not callable. + """ + if not callable(transform): + raise TypeError("loc_transform(transform): transform must be callable") + + token = _LOC_TRANSFORMS.set(_LOC_TRANSFORMS.get() + (transform,)) + try: + yield + finally: + _LOC_TRANSFORMS.reset(token) + + +def _apply_loc_transforms(loc: Any) -> Any: + """Apply the active scoped loc transforms to ``loc``.""" + for transform in reversed(_LOC_TRANSFORMS.get()): + loc = transform(loc) + return loc + + def _set_include_lib_frame(enable: bool) -> None: """Set whether ops are attributed to the closest (library) frame. @@ -392,6 +443,9 @@ def wrapper(*args: Any, **kwargs: Any) -> Any: # ``_mutable_ref`` (see pyir_runtime._pyir_auto_load_arg). log().debug("[dsl_user_op] %s called with %d args", opFunc.__name__, len(args)) args = tuple(map(_pyir_auto_load_arg, args)) + # Value-carrying kwargs escape loops like positional args; auto-load them + # so a post-loop kwarg use reads from the ref, not a stale in-loop SSA. + kwargs = {k: _pyir_auto_load_arg(v) for k, v in kwargs.items()} # Pop loc= from kwargs so callers that still pass it don't break. # The wrapper replaces it only when source-location tracking is enabled. loc: Any = kwargs.pop("loc", None) @@ -408,17 +462,16 @@ def wrapper(*args: Any, **kwargs: Any) -> Any: # outside a kernel). Proceed with loc=None so that the # wrapped function's own validation can still fire. pass + loc = _apply_loc_transforms(loc) # __init__ wrappers either wrap an existing ir.Value (no new # ops, e.g. with_signedness) or build a trivial arith constant # that always verifies. They're also called *very* frequently - # from hot paths like signedness coercion, so we skip both the - # dominance check and the block-diff verify for them — + # from hot paths like signedness coercion, so we skip the + # block-diff verify for them — # materializing `block.operations` on every call would be O(N) # per call and turn kernel build into O(N^2). is_init = getattr(opFunc, "__name__", "") == "__init__" - if not is_init: - _check_operand_dominance(args, kwargs, frameInfo, opFunc) # Snapshot the current insertion block so we can verify newly-built # ops after opFunc returns. The dialect Python bindings strip @@ -515,6 +568,10 @@ def wrapper(*args: Any, **kwargs: Any) -> Any: return res_or_list + # F-COVER: a DSL op wrapper runs natively inside a trace by design; the + # decorator declares that so entry attestation never mistakes it for an + # un-instrumented user function. + wrapper.__pyir_native__ = True # type: ignore[attr-defined] return wrapper diff --git a/python/CuTeDSL/cutlass/__init__.py b/python/CuTeDSL/cutlass/__init__.py index 82a951d7d7..1ed35bf4cc 100644 --- a/python/CuTeDSL/cutlass/__init__.py +++ b/python/CuTeDSL/cutlass/__init__.py @@ -38,11 +38,14 @@ def _ensure_mlir_type_compat() -> None: __version__ = "@CUTLASS_IR_WHEEL_RELEASE_VERSION@" # Monkey patch CUDA version query function from ._mlir._mlir_libs._cutlass_ir._base_dsl import ( + NumericCast as _NumericCast, get_cuda_version as _get_cuda_version, ) from .base_dsl import common as _common +from .base_dsl import typing as _base_dsl_typing _common._get_cuda_version = _get_cuda_version +_base_dsl_typing._native_numeric_cast = _NumericCast # Import CUDA version from base_dsl from .base_dsl.version_info import CUDA_VERSION diff --git a/python/CuTeDSL/cutlass/base_dsl/ast_helpers.py b/python/CuTeDSL/cutlass/base_dsl/ast_helpers.py index d72ec8d57a..55f36ba804 100644 --- a/python/CuTeDSL/cutlass/base_dsl/ast_helpers.py +++ b/python/CuTeDSL/cutlass/base_dsl/ast_helpers.py @@ -31,6 +31,15 @@ from .common import DSLUserCodeError as DSLUserCodeError # star-import re-export from .diagnostics import DiagId +# Generated code names pyir_runtime symbols through the module alias the +# rewriter binds in the function preamble; this import is for ast_helpers' +# OWN use (if_selector witnesses trace-time predicate folds), not a re-export. +from .pyir_runtime import _pyir_witness_predicate_fold +from .multi_stage_manager import ( # noqa: F401 (re-exported via wildcard) + enter_constexpr_loop, + exit_constexpr_loop, +) + class Executor: """ @@ -259,7 +268,9 @@ def ir_loop(func: Callable[..., Any]) -> Any: def if_selector(pred: Any, write_args: list[Any] = []) -> Callable[..., Any]: log().debug("pred [%s] write_args [%s]", pred, write_args) - # Handle Numeric types here? + # Witness a trace-time fold of a watched META predicate so a later staged + # write of its source places refuses loudly. No-op for other predicates. + _pyir_witness_predicate_fold(pred) from .typing import Numeric @@ -460,6 +471,20 @@ def assert_executor(test: Any, msg: str | None = None) -> None: ) +def bool_short_circuits(value: Any, short_circuit_value: bool) -> bool: + """True when the and/or LHS *value* short-circuits in Python: it is a + Python-truth bool -- a plain ``bool``, or a wrapper DECLARING a bool + payload through ``_pyir_raw_payload`` -- whose truth equals + *short_circuit_value*. A staged or non-bool value answers False, so the + rewrite's other arm evaluates (the ``and_``/``or_`` helper).""" + if type(value) is bool: + return value == short_circuit_value + payload = getattr(value, "_pyir_raw_payload", None) + if type(payload) is bool: + return payload == short_circuit_value + return False + + def bool_cast(value: Any) -> bool: if executor._is_dynamic_expression(value): # type: ignore[misc] raise DSLUserCodeError( @@ -487,7 +512,31 @@ def compare_executor(left: Any, comparators: list[Any], ops: list[Any]) -> Any: assert executor._compare_executor is not None, ( "Function must be set before execution." ) - return executor._compare_executor(left, comparators, ops) + if "is" in ops or "is not" in ops: + # Identity legs over tracker-minted wrappers refuse-or-witness + # (LangRef 3.12 section 6.10.3); inert outside a PyIR trace scope. + from .pyir_core import _pyir_identity_compare_choke + + _pyir_identity_compare_choke(left, comparators, ops) + result = executor._compare_executor(left, comparators, ops) + # Gate-once evidence (see pyir_state): record only the provable membership-gate + # shape (one ``not in`` over a plain set); clear on every other comparison. + try: + from .pyir_runtime import _PYIR_LAST_NOTIN_COMPARE + + if ( + len(ops) == 1 + and ops[0] == "not in" + and type(result) is bool + and type(left) in (str, int, bool) + and type(comparators[0]) is set + ): + _PYIR_LAST_NOTIN_COMPARE[0] = (comparators[0], left, result) + else: + _PYIR_LAST_NOTIN_COMPARE[0] = None + except Exception: + pass + return result # ============================================================================= @@ -666,83 +715,6 @@ def closure_check( ) -def _is_dsl_traced_callable(fn: object) -> bool: - """Whether *fn* is a jit-/kernel-decorated wrapper. - - Such a wrapper re-traces its body on every staged pass (the jit decorator - tags the wrapper with ``_dsl_cls`` / ``_dsl_object``), so its effects are - NOT frozen at the first trace. - """ - return hasattr(fn, "_dsl_cls") or hasattr(fn, "_dsl_object") - - -_REGION_BUILDER_SCOPES = ( - "loop_body_", - "then_block_", - "else_block_", - "if_region_", - "while_region_", - "while_before_block_", - "while_after_block_", - "ifexp_then_block_", - "ifexp_else_block_", -) - - -def lambda_capture_check(candidates: list[Any]) -> None: - """Reject a ``lambda`` invoked inside staged CF that captures an enclosing - local. - - A bare ``lambda`` is invoked as raw Python at trace time; its body reads - captured variables through a closure cell bound to the ENCLOSING function's - local -- not the staged region's loop-carried slot. A captured meta that - is mutated inside the region therefore reads its first-pass value forever - (a silent miscompile). Only lambdas are handled here; ordinary (``def``) - closures are out of scope. Jit-decorated lambdas trace per pass, - and lambdas capturing nothing (or only globals / other functions) are safe. - - *candidates* is the region's bare-name calls that resolve to values in - scope (rather than tracked callables); non-lambdas are skipped. Emitted - only by the PyIR preprocessor subclass, so non-pyir compilation never runs - this check. - """ - for fn in candidates: - if not isinstance(fn, types.FunctionType): - continue - if _is_dsl_traced_callable(fn): - continue - if getattr(fn, "__name__", None) != "": - continue - # A lambda DEFINED INSIDE the staged region cannot go stale: the - # region builder re-creates it on the trace pass, so its closure - # cells hold that pass's bindings. Only lambdas defined OUTSIDE and - # called INSIDE freeze their captures. Region-builder scopes are - # identifiable from the generated block names in the qualname. - qualname = getattr(fn, "__qualname__", "") - if any(f".{part}" in qualname for part in _REGION_BUILDER_SCOPES): - continue - # Read the closure captures natively. A lambda's captured (nonlocal) - # names are ``co_freevars`` positionally paired with its ``__closure__`` - # cells; globals/builtins are not cells, so they are excluded. An - # UNFILLED cell cannot be vouched for -> map it to ``None`` so it takes - # the reject path below (fail closed). - for name, cell in zip(fn.__code__.co_freevars, fn.__closure__ or ()): - try: - value = cell.cell_contents - except ValueError: - value = None - if value is not None and ( - inspect.ismodule(value) - or inspect.isfunction(value) - or inspect.ismethod(value) - ): - continue - raise DSLUserCodeError( - DiagId.SCOPE_LAMBDA_CAPTURE, - var_name=name, - ) - - @dataclass class FormattedValue: """ diff --git a/python/CuTeDSL/cutlass/base_dsl/ast_preprocessor.py b/python/CuTeDSL/cutlass/base_dsl/ast_preprocessor.py index d88dd71bdb..745a505190 100644 --- a/python/CuTeDSL/cutlass/base_dsl/ast_preprocessor.py +++ b/python/CuTeDSL/cutlass/base_dsl/ast_preprocessor.py @@ -120,43 +120,6 @@ def intersections(self, others: list[set[str]]) -> "OrderedSet": return result -@dataclass -class ImportInfo: - """ - Information about an import expression. - """ - - module_path: str - attr_name: str | None - alias_name: str - - -@dataclass -class TryImportInfo: - """ - Represents information about a try-import block in the AST. - - This dataclass is used to capture and organize the import statements that appear - within the different clauses of a try-except-else-finally block. Each field holds - a list of import statements (or related nodes) that are encountered in the corresponding - clause of the try block. - - Attributes: - try_imports (list): Import statements found in the 'try' clause. - except_imports (list): Import statements found in any 'except' clauses. - else_imports (list): Import statements found in the 'else' clause, if present. - finally_imports (list): Import statements found in the 'finally' clause, if present. - - This structure allows the preprocessor to track and process imports that are conditionally - executed depending on exception handling logic. - """ - - try_imports: "list[ImportInfo | TryImportInfo]" - except_imports: "list[ImportInfo | TryImportInfo]" - else_imports: "list[ImportInfo | TryImportInfo]" - finally_imports: "list[ImportInfo | TryImportInfo]" - - @dataclass class ScopeManager: """ @@ -400,6 +363,13 @@ def set_location( return result +def _mark_synth_scope(func_def: ast.FunctionDef) -> ast.FunctionDef: + """SYNTHESIZED-scope fact: this def is a rewrite-created arm/body block, + not a source scope; the call-boundary pass keys frame reflection on it.""" + func_def._pyir_synth_scope = True # type: ignore[attr-defined] + return func_def + + _ComprehensionT = TypeVar( "_ComprehensionT", ast.ListComp, ast.SetComp, ast.GeneratorExp, ast.DictComp ) @@ -418,6 +388,10 @@ class DSLPreprocessor(ast.NodeTransformer): DECORATOR_IF_STATEMENT = "if_selector" DECORATOR_WHILE_STATEMENT = "while_selector" IF_EXECUTOR = "if_executor" + + # F-COVER: the base rewrite emits no PyIR choke set, so functions it + # rewrites carry no ``__pyir_rewritten__`` attestation stamp. + choke_set_version: "int | None" = None IFEXP_EXECUTOR = "ifExp_executor" WHILE_EXECUTOR = "while_executor" ASSERT_EXECUTOR = "assert_executor" @@ -579,6 +553,12 @@ def _inject_default_arg_values( if source_name is not None and source_name not in exec_globals: exec_globals[source_name] = default_val + def _post_visit_function_body(self, func_def: ast.FunctionDef) -> "list[ast.stmt]": + """Hook for a mode-specific post-pass over the fully instrumented + function body; returns preamble statements to bind right after the + module-alias imports. The base rewrite has none.""" + return [] + def transform_function( self, func_name: str, function_pointer: Callable[..., Any] ) -> list[ast.stmt]: @@ -650,9 +630,16 @@ def transform_function( # Step 2. Transform the function transformed_tree = self.visit(tree) + # Step 2.5. Mode-specific post-pass over the instrumented tree (the + # body is fully visited at this point); any returned preamble + # statements bind after the module-alias imports below. + preamble_stmts: "list[ast.stmt]" = [] + if isinstance(transformed_tree.body[0], ast.FunctionDef): + preamble_stmts = self._post_visit_function_body(transformed_tree.body[0]) + # Step 3. Import cutlass and base_dsl top_module_name = ".".join(self.client_module_name) - import_stmts = [] + import_stmts: "list[ast.stmt]" = [] if self.session_data.import_top_module: import_stmts.append( ast.Import( @@ -669,6 +656,7 @@ def transform_function( assert len(transformed_tree.body) == 1 assert isinstance(transformed_tree.body[0], ast.FunctionDef) + import_stmts.extend(preamble_stmts) transformed_tree.body[0].body = import_stmts + transformed_tree.body[0].body # Remove all decorators from top level function transformed_tree.body[0].decorator_list = [] @@ -860,7 +848,7 @@ def analyze_region_variables( node: ast.For | ast.If | ast.While, active_symbols: list[set[str]], active_callables: list[set[str]], - ) -> tuple[list[str], int, list[str], list[str]]: + ) -> tuple[list[str], int, list[str]]: """ Analyze loop-carried and closure variables in a control-flow region. @@ -958,21 +946,10 @@ def visit_Call(self, node: ast.Call) -> None: for name in called_functions.intersections(active_callables) if name != current_function_name ] - # Bare-name calls that resolve to an in-scope value but NOT a tracked - # callable (``x = lambda ...``, ``x = factory()``). A raw lambda among - # these, invoked inside staged CF, reads its captures from the enclosing - # frame rather than the region's loop-carried slots and would silently - # miscompile; the PyIR subclass emits ``lambda_capture_check`` to reject - # that case (non-lambdas are skipped there; the base emits nothing). - called_value_symbols: list[str] = list( - called_functions.intersections(active_symbols) - - OrderedSet(called_functions_list) - ) return ( write_args_list + invoked_args_list, len(write_args_list), called_functions_list, - called_value_symbols, ) def extract_range_args( @@ -1145,17 +1122,19 @@ def create_loop_function( ) return ast.copy_location( - ast.FunctionDef( - name=func_name, - args=ast.arguments( - posonlyargs=[], - args=func_args, - kwonlyargs=[], - kw_defaults=[], - defaults=[], - ), - body=transformed_body, - decorator_list=[decorator], + _mark_synth_scope( + ast.FunctionDef( + name=func_name, + args=ast.arguments( + posonlyargs=[], + args=func_args, + kwonlyargs=[], + kw_defaults=[], + defaults=[], + ), + body=transformed_body, + decorator_list=[decorator], + ) ), node, ) @@ -1538,23 +1517,13 @@ def _create_closure_check_call( ) ) - def _create_lambda_check_call( - self, called_value_symbols: list[str], node: ast.stmt - ) -> ast.Expr | None: - """Base stub: emit nothing. - - The lambda-capture guard is a silent-miscompile only under staged - tracing, so it is emitted ONLY by the PyIR preprocessor subclass' - override. The base returns None -> non-pyir compilation emits a - byte-identical region with no ``lambda_capture_check`` call. - """ - return None - - def _prepare_loop_induction_var(self, node: ast.For) -> None: - """Prepare loop induction variable before function creation. - - Override for custom behavior (e.g., mark variable for special handling). - """ + def _prepare_loop_induction_var( + self, + node: ast.For, + target_is_live_after_loop: bool = False, + loop_carried_var_name: str | None = None, + ) -> None: + """Prepare loop induction variable before function creation.""" pass # No preparation needed in base class def _cleanup_loop_induction_var(self, node: ast.For) -> None: @@ -1617,11 +1586,16 @@ def transform_for_loop( assert isinstance(node.iter, ast.Call) start_expr, stop_expr, step_expr, has_step = self.extract_range_args(node.iter) + # Template method: a derived class may instrument the bound expressions, + # consumed at loop-selector call time outside the visited body. + start_expr = self._prepare_loop_bound_expr(start_expr) + stop_expr = self._prepare_loop_bound_expr(stop_expr) + step_expr = self._prepare_loop_bound_expr(step_expr) unroll, unroll_full = self.extract_unroll_args(node.iter) prefetch_stages = self.extract_prefetch_stages_args(node.iter) vectorize = self.extract_vectorize_args(node.iter) at_least_once = self.extract_at_least_once_args(node.iter) - write_args, full_write_args_count, called_closures, called_value_symbols = ( + write_args, full_write_args_count, called_closures = ( self.analyze_region_variables(node, active_symbols, active_callables) ) @@ -1654,15 +1628,18 @@ def transform_for_loop( if cc is not None: exprs.append(cc) - lc = self._create_lambda_check_call(called_value_symbols, node) - if lc is not None: - exprs.append(lc) - func_name = f"loop_body_{self.session_data.counter}" self.session_data.counter += 1 - # Template method: prepare induction variable (e.g., mark for special handling) - self._prepare_loop_induction_var(node) + # Template method: prepare induction variable (e.g., mark for special + # handling); passes the live-out signal + synthetic carry name. + self._prepare_loop_induction_var( + node, + target_is_live_after_loop=target_var_is_active_before_loop, + loop_carried_var_name=( + loop_carried_var_name if target_var_is_active_before_loop else None + ), + ) func_def = self.create_loop_function( func_name, @@ -1817,8 +1794,15 @@ def visit_Call(self, node: ast.Call) -> ast.Call: if isinstance(func, ast.Name): # AST rewrite only redirect call to bool to bool_cast # If `bool` escapes as a symbol, usually it means type check, do not rewrite it - if func.id == "bool": - return ast.copy_location( + # Any other call shape has no `bool_cast` spelling, so it is left as a + # plain `bool` call for Python itself to accept or reject + if ( + func.id == "bool" + and len(node.args) == 1 + and node.keywords == [] + and not isinstance(node.args[0], ast.Starred) + ): + redirected = ast.copy_location( ast.Call( func=ast.Call( func=_create_module_attribute( @@ -1834,6 +1818,10 @@ def visit_Call(self, node: ast.Call) -> ast.Call: ), node, ) + # Machinery application of the redirector's result: the outer + # call is not a user call, so the boundary pass skips it. + redirected._pyir_synth = True # type: ignore[attr-defined] + return redirected elif func.id == "super" and node.args == [] and node.keywords == []: # If it's a Python3 argument free super(), rewrite to old style super with args # So if this call is under dynamic control flow, it still works. @@ -1926,7 +1914,7 @@ def visit_AugAssign(self, node: ast.AugAssign) -> ast.AugAssign | list[ast.stmt] self.generic_visit(node) return node - def visit_AnnAssign(self, node: ast.AnnAssign) -> ast.AnnAssign: + def visit_AnnAssign(self, node: ast.AnnAssign) -> "ast.stmt | list[ast.stmt]": self._visit_target(node.target) self.generic_visit(node) return node @@ -1964,12 +1952,9 @@ def get_dsl_decorator_index(self, decorator_list: list[ast.expr]) -> Any: # "attributes" passes kernel function attributes (e.g. launch bounds) # to the compiler — its presence shouldn't prevent the preprocessor # from recognizing @cute.kernel(attributes=...) as a DSL decorator. - # "is_experimental" routes the function through the experimental - # CuTe DSL (see ``CuTeDSL.jit`` / ``CuTeDSL.kernel``). _known_dsl_kwargs = { "preprocess", "attributes", - "is_experimental", } for i, d in enumerate(decorator_list): @@ -2091,7 +2076,9 @@ def visit_Nonlocal(self, node: ast.Nonlocal) -> ast.Nonlocal: self.generic_visit(node) return node - def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef: + def visit_FunctionDef( + self, node: ast.FunctionDef + ) -> "ast.FunctionDef | list[ast.stmt]": # Add self to active symbols of parent scope self.session_data.scope_manager.add_to_callables(node.name) @@ -2170,7 +2157,7 @@ def visit_While(self, node: ast.While) -> ast.While | list[ast.stmt]: with self.session_data.scope_manager.enter_control_flow_scope(): self.check_early_exit(node, "while") - write_args, full_write_args_count, called_closures, called_value_symbols = ( + write_args, full_write_args_count, called_closures = ( self.analyze_region_variables(node, active_symbols, active_callables) ) exprs = [] @@ -2179,10 +2166,6 @@ def visit_While(self, node: ast.While) -> ast.While | list[ast.stmt]: if cc is not None: exprs.append(cc) - lc = self._create_lambda_check_call(called_value_symbols, node) - if lc is not None: - exprs.append(lc) - func_name = f"while_region_{self.session_data.counter}" self.session_data.counter += 1 @@ -2369,7 +2352,7 @@ def visit_IfExp(self, node: ast.IfExp) -> ast.Call: # Insert the block definitions into the most recent (innermost) region before the statement self.session_data.region_stack[-1].append_new_stmts( - [then_block_def, else_block_def] + [_mark_synth_scope(then_block_def), _mark_synth_scope(else_block_def)] ) # Create the executor call node, wiring up the predicate and newly synthesized blocks @@ -2492,7 +2475,7 @@ def visit_If(self, node: ast.If) -> ast.If | list[ast.stmt]: with self.session_data.scope_manager.enter_control_flow_scope(): self.check_early_exit(node, "if") - yield_args, full_write_args_count, called_closures, called_value_symbols = ( + yield_args, full_write_args_count, called_closures = ( self.analyze_region_variables(node, active_symbols, active_callables) ) exprs = [] @@ -2501,10 +2484,6 @@ def visit_If(self, node: ast.If) -> ast.If | list[ast.stmt]: if cc is not None: exprs.append(cc) - lc = self._create_lambda_check_call(called_value_symbols, node) - if lc is not None: - exprs.append(lc) - func_name = f"if_region_{self.session_data.counter}" self.session_data.counter += 1 @@ -2516,12 +2495,16 @@ def visit_If(self, node: ast.If) -> ast.If | list[ast.stmt]: return exprs + [func_def] + assign def generate_get_locals_or_none_call(self, write_args: list[str]) -> ast.Call: + # The inner ``locals()`` is threading plumbing, not a user frame + # reflection: mark it so the call-boundary pass leaves it bare. + plumbing_locals = ast.Call( + func=ast.Name(id="locals", ctx=ast.Load()), args=[], keywords=[] + ) + plumbing_locals._pyir_synth = True # type: ignore[attr-defined] return ast.Call( func=_create_module_attribute("get_locals_or_none"), args=[ - ast.Call( - func=ast.Name(id="locals", ctx=ast.Load()), args=[], keywords=[] - ), + plumbing_locals, ast.List( elts=[ast.Constant(value=arg) for arg in write_args], ctx=ast.Load(), @@ -2584,11 +2567,13 @@ def create_if_function( # Create then block then_block = ast.copy_location( - ast.FunctionDef( - name=then_block_name, - args=func_then_else_arguments, - body=then_body + [ast.Return(value=return_list)], - decorator_list=[], + _mark_synth_scope( + ast.FunctionDef( + name=then_block_name, + args=func_then_else_arguments, + body=then_body + [ast.Return(value=return_list)], + decorator_list=[], + ) ), node, ) @@ -2671,12 +2656,20 @@ def create_if_function( # else: # if pred: # And under both cases, the `pred` can be a const_expr, so we need to handle it here. + # Statements hoisted while visiting the elif test (walrus + # lowering, read/effect anchors) must execute only when + # every earlier arm's condition is false: collect them + # into the synthesized else block, not the region + # enclosing the whole if. + elif_pre: list[ast.stmt] = [] if self.is_node_constexpr(elif_node): - check = self._handle_constexpr_elif(elif_node) + with Region(self.session_data, new_value=elif_pre): + check = self._handle_constexpr_elif(elif_node) else_block = ast.FunctionDef( name=else_block_name, args=func_then_else_arguments, - body=[ + body=elif_pre + + [ check, elif_node, ast.Return(value=return_list), @@ -2685,13 +2678,18 @@ def create_if_function( ) else: # Recursion for nested elif - nested_if = self.create_if_function( - nested_if_name, elif_node, write_args, full_write_args_count - ) + with Region(self.session_data, new_value=elif_pre): + nested_if = self.create_if_function( + nested_if_name, + elif_node, + write_args, + full_write_args_count, + ) else_block = ast.FunctionDef( name=else_block_name, args=func_then_else_arguments, - body=[ + body=elif_pre + + [ nested_if, ast.Return( value=ast.Name(id=nested_if_name, ctx=ast.Load()) @@ -2723,6 +2721,7 @@ def create_if_function( decorator_list=[], ) + _mark_synth_scope(else_block) # Add else_block to execute keywords execute_keywords.append( ast.keyword( @@ -2769,6 +2768,11 @@ def _prepare_while_condition_vars( """ return [] # No preparation needed in base class + def _prepare_loop_bound_expr(self, expr: ast.expr) -> ast.expr: + """Instrument a staged-for range bound expression, evaluated at + loop-selector call time outside the visited body; base is identity.""" + return expr + def _prepare_loop_body_vars( self, node: "ast.For | ast.While", @@ -2882,11 +2886,13 @@ def while_after_block(*write_args): ) while_before_stmts.append(ast.Return(value=while_before_return_list)) while_before_block = ast.copy_location( - ast.FunctionDef( - name=while_before_block_name, - args=block_args, - body=while_before_stmts, - decorator_list=[], + _mark_synth_scope( + ast.FunctionDef( + name=while_before_block_name, + args=block_args, + body=while_before_stmts, + decorator_list=[], + ) ), test_expr, ) @@ -2907,11 +2913,13 @@ def while_after_block(*write_args): while_after_stmts.append(ast.Return(value=yield_args_ast_name_list)) while_after_block = ast.copy_location( - ast.FunctionDef( - name=while_after_block_name, - args=block_args, - body=while_after_stmts, - decorator_list=[], + _mark_synth_scope( + ast.FunctionDef( + name=while_after_block_name, + args=block_args, + body=while_after_stmts, + decorator_list=[], + ) ), node, ) diff --git a/python/CuTeDSL/cutlass/base_dsl/common.py b/python/CuTeDSL/cutlass/base_dsl/common.py index 455371efd6..ab2449d84d 100644 --- a/python/CuTeDSL/cutlass/base_dsl/common.py +++ b/python/CuTeDSL/cutlass/base_dsl/common.py @@ -38,12 +38,43 @@ "active_env_manager", default=None ) -_CUDA_INVALID_LAUNCH_VALUE_ERRORS = { - "CUDA_ERROR_INVALID_VALUE", - "cudaErrorInvalidValue", +_CUDA_ERROR_NAME_ALIASES = { + "cudaErrorNoKernelImageForDevice": "CUDA_ERROR_NO_BINARY_FOR_GPU", + "cudaErrorMemoryAllocation": "CUDA_ERROR_OUT_OF_MEMORY", + "cudaErrorInitializationError": "CUDA_ERROR_NOT_INITIALIZED", } +def _cuda_error_lookup_key(error_name: str) -> str: + """Fold the driver and runtime spellings of one CUDA error onto one key. + + The same failure reaches this module as ``CUDA_ERROR_INVALID_VALUE`` or as + ``cudaErrorInvalidValue`` depending on which CUDA API reported it, so tables + keyed by error name must not depend on the spelling. Pairs that differ by + more than case and punctuation are folded through + ``_CUDA_ERROR_NAME_ALIASES`` first. + """ + canonical = _CUDA_ERROR_NAME_ALIASES.get(error_name, error_name).lower() + for prefix in ("cuda_error_", "cudaerror"): + if canonical.startswith(prefix): + canonical = canonical[len(prefix) :] + break + return canonical.replace("_", "") + + +def _lookup_by_cuda_error( + table: Dict[str, Any], lookup_key: str, default: Any = "" +) -> Any: + """Look up a table keyed by CUDA error name, ignoring the spelling.""" + for name, value in table.items(): + if _cuda_error_lookup_key(name) == lookup_key: + return value + return default + + +_CUDA_INVALID_LAUNCH_VALUE_KEY = _cuda_error_lookup_key("CUDA_ERROR_INVALID_VALUE") + + def get_current_env_manager() -> Any: """Return the env manager for the active DSL context, if any. @@ -71,53 +102,6 @@ def active_env_manager(env_manager: Any) -> Generator[None, None, None]: _active_env_manager.reset(token) -# The DSL package's own root as imported (wheel: site-packages/cutlass; dev: the -# build python_packages farm -- modules imported through the package carry THIS -# path in ``co_filename`` even though the files are symlinks into the source -# tree). Deliberately NOT a resolved-source-tree anchor: the source repo also -# contains tests/examples, which are AUTHOR code. -_DSL_PKG_ROOT: str = ( - os.path.dirname(os.path.dirname(os.path.abspath(__file__))) + os.sep -) -# Top-level name of the resolved source layout (dev checkouts import some -# internal modules directly from the source tree top package). Derived at -# runtime from this module's resolved location -- no hardcoded layout name. -_DSL_RESOLVED_TOP: str = os.path.basename( - os.path.dirname(os.path.dirname(os.path.realpath(__file__))) -) - - -def is_dsl_internal_code(filename: Optional[str], module_name: str = "") -> bool: - """Whether code at *filename* (module *module_name*) is DSL-internal or - standard-library -- i.e. NOT the DSL user's own code. - - THE single classifier for "ours vs author code" decisions (the - context-manager / lambda staging guards, ...). DSL code is recognized by the - package's own path (as imported) or, for dev-layout direct imports, by the - resolved top-level package NAME; the stdlib by top-level module name (the - official ``sys.stdlib_module_names`` API). Files that merely LIVE in the - source repo (tests, examples, user scripts run as ``__main__``) classify as - author code. ``filename=None`` (no code object, e.g. builtins / C - extensions) classifies via the module's ``__file__`` when importable. - """ - top = (module_name or "").split(".", 1)[0] - if top and top in getattr(sys, "stdlib_module_names", frozenset()): - return True - if top and top == _DSL_RESOLVED_TOP: - return True - if filename is None: - # No source: builtins / C extensions. Resolve through the module's - # ``__file__`` when importable so DSL-shipped extensions classify as - # internal; unresolvable means we cannot vouch for it. - mod = sys.modules.get(module_name) - filename = getattr(mod, "__file__", "") or "" - if not filename: - return False - if filename.startswith("<"): - return True # exec/ frames: generated glue, not author files - return os.path.abspath(filename).startswith(_DSL_PKG_ROOT) - - def is_pyir_enabled() -> bool: """Return True when the user has set ``ENABLE_PYIR=True``.""" env_manager = get_current_env_manager() @@ -280,14 +264,15 @@ class DSLRuntimeError(DSLBaseError): _ARCH_RELATED_CUDA_ERRORS = frozenset( - { + _cuda_error_lookup_key(name) + for name in ( "CUDA_ERROR_INVALID_SOURCE", "CUDA_ERROR_NO_BINARY_FOR_GPU", "CUDA_ERROR_INVALID_PTX", "CUDA_ERROR_UNSUPPORTED_PTX_VERSION", "CUDA_ERROR_NO_DEVICE", "CUDA_ERROR_INVALID_DEVICE", - } + ) ) @@ -311,17 +296,18 @@ def _get_friendly_cuda_error_message( from .runtime.cuda import get_device_info error_name = _normalize_cuda_error_name(error_name) + lookup_key = _cuda_error_lookup_key(error_name) env_manager = get_current_env_manager() target_arch = env_manager.arch if env_manager is not None else "unknown" - arch_is_relevant = error_name in _ARCH_RELATED_CUDA_ERRORS + arch_is_relevant = lookup_key in _ARCH_RELATED_CUDA_ERRORS invalid_launch_value_suggestion = ( "Check `.launch(...)`: grid, block, dynamic shared memory, stream, " "and attributes. Keep block.x * block.y * block.z <= " "maxThreadsPerBlock. If only threadIdx.x is used, launch with " "block=(threads, 1, 1)." ) - if error_name in _CUDA_INVALID_LAUNCH_VALUE_ERRORS: + if lookup_key == _CUDA_INVALID_LAUNCH_VALUE_KEY: message = f"CUDA launch failed: {error_name} ({error_code})" return message, "", invalid_launch_value_suggestion @@ -388,13 +374,13 @@ def _get_friendly_cuda_error_message( ), "cudaErrorInsufficientDriver": ( "1. Run nvidia-smi to confirm CUDA driver version", - "2. Ensure the CUDA driver version meets the requirement of the installed cuda-python package", + "2. Ensure the CUDA driver version meets the requirement of the installed cuda-bindings package", ), } message = ( f"{error_name} (error code: {error_code}) \n" - f"{additional_info.get(error_name, '')} \n\n{Colors.RESET}" + f"{_lookup_by_cuda_error(additional_info, lookup_key)} \n\n{Colors.RESET}" ) # Add debug information @@ -444,7 +430,7 @@ def _get_friendly_cuda_error_message( f"\n{Colors.YELLOW}ℹ️ Could not retrieve GPU info: {str(e)}{Colors.RESET}" ) - return message, debug_info, error_suggestions.get(error_name, "") + return message, debug_info, _lookup_by_cuda_error(error_suggestions, lookup_key) class DSLCudaRuntimeError(DSLBaseError): @@ -465,7 +451,8 @@ def __init__( self._error_name = error_name normalized_error_name = _normalize_cuda_error_name(error_name) concise_launch_error = ( - normalized_error_name in _CUDA_INVALID_LAUNCH_VALUE_ERRORS + _cuda_error_lookup_key(normalized_error_name) + == _CUDA_INVALID_LAUNCH_VALUE_KEY ) if concise_launch_error: self.code = "CUDA_LAUNCH_INVALID_CONFIG" @@ -752,6 +739,7 @@ def __init__( current_frame = inspect.currentframe() frame = current_frame.f_back if current_frame else None frameInfo = inspect.getframeinfo(frame) if frame else None + del current_frame # Try to translate MLIR/nanobind errors if no custom message provided self.original_error = str(message) diff --git a/python/CuTeDSL/cutlass/base_dsl/compile_backend.py b/python/CuTeDSL/cutlass/base_dsl/compile_backend.py index c4ccf3b145..fb7cab6ae4 100644 --- a/python/CuTeDSL/cutlass/base_dsl/compile_backend.py +++ b/python/CuTeDSL/cutlass/base_dsl/compile_backend.py @@ -98,7 +98,7 @@ def _resolve_compile_cache(self, ctx: CompileContext) -> CompileCacheState: dsl = self._dsl load_from_file_cache = False - if ctx.cache_enabled: + if ctx.cache_enabled and not dsl.envar.disable_file_caching: assert ctx.module_hash is not None fn = load_cache_from_path( dsl.name, ctx.module_hash, bytecode_reader=read_bytecode_and_check_crc32 diff --git a/python/CuTeDSL/cutlass/base_dsl/compiler.py b/python/CuTeDSL/cutlass/base_dsl/compiler.py index a96363f280..62e447d173 100644 --- a/python/CuTeDSL/cutlass/base_dsl/compiler.py +++ b/python/CuTeDSL/cutlass/base_dsl/compiler.py @@ -30,7 +30,6 @@ from . import diagnostics as _diagnostics from .utils.logger import log from .env_manager import EnvironmentVarManager - _SCRIPT_PATH = os.path.dirname(os.path.abspath(__file__)) @@ -47,6 +46,59 @@ def _split_options(text: str) -> list: return shlex.split(text) +# msvcrt's FILE is 48 bytes on x64 and stdout is __iob[1]; _IONBF is 4. +_MSVCRT_FILE_SIZE = 48 +_IONBF = 4 +_jit_c_stdout_unbuffered = False + + +def _unbuffer_jit_c_stdout() -> None: + """Unbuffer the C stdout that the JIT's ``printf`` writes to (Windows). + + LLVM binds the JIT's ``printf`` by scanning the process's loaded modules + for an export of that name. The UCRT does not export one -- ``printf`` is + inline in its -- so the search lands on the legacy + ``msvcrt.dll``, whose C stdout is a stream of its own, unrelated to the one + Python writes through. On a pipe it is fully buffered and flushed only when + msvcrt runs its exit handler, so every ``cute.printf`` line surfaces after + all Python output rather than interleaving in program order. + + lit already pins the Python side with ``PYTHONUNBUFFERED=1``; this pins the + C side to match. Done once when the execution engine is built, so nothing + lands on the execution path, and it touches only msvcrt's stdout, which no + other part of the process writes to. No-op off Windows, where ``printf`` + resolves to the same libc Python is already sharing. + """ + global _jit_c_stdout_unbuffered + if _jit_c_stdout_unbuffered or os.name != "nt": + return + # Only ever attempt this once, successful or not. + _jit_c_stdout_unbuffered = True + import ctypes + + try: + msvcrt = ctypes.CDLL("msvcrt.dll") + iob_func = getattr(msvcrt, "__iob_func") + iob_func.restype = ctypes.c_void_p + msvcrt._fileno.restype = ctypes.c_int + msvcrt._fileno.argtypes = [ctypes.c_void_p] + msvcrt.setvbuf.argtypes = [ + ctypes.c_void_p, + ctypes.c_void_p, + ctypes.c_int, + ctypes.c_size_t, + ] + stdout = iob_func() + _MSVCRT_FILE_SIZE + # Confirm the FILE stride against the descriptor rather than trusting + # the struct layout blindly; leave the stream alone if it disagrees. + if msvcrt._fileno(stdout) != 1: + log().warning("msvcrt stdout not found at __iob[1]; leaving buffered") + return + msvcrt.setvbuf(stdout, None, _IONBF, 0) + except (OSError, AttributeError) as e: + log().warning(f"could not unbuffer the JIT's C stdout: {e}") + + sys.path.append(_SCRIPT_PATH) from .._mlir import ir @@ -266,6 +318,7 @@ def jit( shared_libs: collections.abc.Sequence[str] = (), ) -> Any: """Wraps the module in a JIT execution engine.""" + _unbuffer_jit_c_stdout() return self.execution_engine.ExecutionEngine( module, opt_level=opt_level, shared_libs=shared_libs ) @@ -730,6 +783,13 @@ class HostTarget(StringCompileOption): Presets:: linux-aarch64 → aarch64-unknown-linux-gnu + qnx8-aarch64 → aarch64-unknown-nto-qnx8.0.0 + + LLVM has no QNX target, so the QNX triple resolves to an unknown OS + and generic AArch64 ELF codegen: the emitted object is byte-identical + to the ``linux-aarch64`` one, and QNX uses the same AAPCS64 ABI. The + distinct triple exists so the AOT runtime library lookup + (``aot_config --target``) can tell the two targets apart. Long form (explicit tuning / escape hatch):: @@ -751,6 +811,7 @@ class HostTarget(StringCompileOption): _PRESETS: "dict[str, tuple[str, str, str]]" = { "linux-aarch64": ("aarch64-unknown-linux-gnu", "", ""), + "qnx8-aarch64": ("aarch64-unknown-nto-qnx8.0.0", "", ""), } def __init__(self, val: str = "") -> None: @@ -1606,6 +1667,13 @@ def _compile(self, func: Any, *args: Any, **kwargs: Any) -> Any: if not hasattr(func, "_dsl_object"): raise DSLUserCodeError(_diagnostics.DiagId.CALL_MISSING_JIT_DECORATOR) + # Reject a @cute.kernel target. + if getattr(func, "_decorator_kind", "jit") == "kernel": + raise DSLUserCodeError( + _diagnostics.DiagId.CALL_KERNEL_TARGET, + function_name=getattr(func, "__name__", ""), + ) + # Validate the migration-aid ``is_experimental`` kwarg against # the routing already baked into the function by its jit/kernel # decorator. This is *not* a behavior switch: both diff --git a/python/CuTeDSL/cutlass/base_dsl/diagnostics.py b/python/CuTeDSL/cutlass/base_dsl/diagnostics.py index 7c055bc985..07f8d52915 100644 --- a/python/CuTeDSL/cutlass/base_dsl/diagnostics.py +++ b/python/CuTeDSL/cutlass/base_dsl/diagnostics.py @@ -733,6 +733,7 @@ def _format_internal_error_diagnostic( # diagnostics so backend failures and internal DSL failures are scannable in # the same way. is_verifier_error = _is_internal_verifier_error(err.message, cause_text) + is_dominance_escape = is_verifier_error and _is_dominance_escape_cause(cause_text) headline = ( "The compiler could not build valid IR for this code." if is_verifier_error @@ -757,6 +758,24 @@ def _format_internal_error_diagnostic( if verifier_detail: summary = f"{summary}: {verifier_detail}" parts.extend(_format_labeled_text("error", summary)) + if is_dominance_escape: + def_frame = _dominance_escape_def_frame(cause_text) + if def_frame: + parts.extend( + _format_labeled_text("note", "the escaping value is created here:") + ) + parts.append(def_frame) + parts.extend( + _format_labeled_text( + "note", + "a staged value created inside a for/while/if body is used " + "outside that body. The tracer threads plain local variables " + "across staged control flow, but it cannot see rebinds made " + "through other channels: object attributes assigned inside " + "helper methods, values stashed in module-level or aliased " + "containers, and closure captures.", + ) + ) else: parts.extend( _format_labeled_text( @@ -771,6 +790,34 @@ def _format_internal_error_diagnostic( parts.extend(_format_internal_cause(cause_text)) if is_verifier_error: + # A mixed *_ENABLE_PYIR configuration is the one declared config fact + # producing this failure shape (a cross-DSL kernel compile leaves + # region-crossing SSA behind): name it before the generic advice. + try: + from .pyir_state import _pyir_mixed_mode_hint + + _mode_hint = _pyir_mixed_mode_hint() + except Exception: + _mode_hint = None + if _mode_hint: + parts.extend(_format_labeled_text("suggestion", _mode_hint)) + if is_dominance_escape: + parts.extend( + _format_labeled_text( + "suggestion", + "Carry the value as a plain local variable rebound directly " + "in the loop or branch body, or re-derive it at the point of " + "use instead of reusing a value captured inside the region.", + ) + ) + if not _PY_LOC_RE.search(cause_text): + parts.extend( + _format_labeled_text( + "suggestion", + "Re-run with CUTE_DSL_LINEINFO=1 to see where the " + "escaping value is created.", + ) + ) parts.extend( _format_labeled_text( "suggestion", @@ -802,6 +849,32 @@ def _is_internal_verifier_error(message: str, cause_text: str) -> bool: ) +def _is_dominance_escape_cause(cause_text: str) -> bool: + """A dominance failure at trace-module verify almost always means a staged + value leaked across region boundaries through a Python-side channel the + tracer does not thread (attribute writes in helper methods, module-level + or aliased containers, closure captures).""" + return "does not dominate this use" in cause_text + + +def _dominance_escape_def_frame(cause_text: str) -> str | None: + """Source frame for the verifier's 'operand defined here' note, available + when the IR carries Python locations (PyIR trace or lineinfo builds).""" + for line in cause_text.splitlines(): + if "operand defined here" not in line: + continue + loc_match = _PY_LOC_RE.search(line) + if not loc_match: + return None + frame = _format_compiler_source_frame( + loc_match.group("file"), + int(loc_match.group("line")), + int(loc_match.group("col")), + ) + return "\n".join(frame) if frame else None + return None + + def _brief_internal_error(message: str) -> str: if "ICE IR Verification Failed" in message: return "IR verification failed" @@ -1510,19 +1583,17 @@ class DiagId(_DiagMixin, enum.Enum): "If it is a fixed setting, set it once before the for/while/if.", ), ) - PHASE_PREDICATE_FOLDED_STALE = ( - "`{var}` (a {meta}, value {value}) already decided {fold_kind} at " - "{fold_location} while your code was being traced, but it is modified at " - "{mut_location}, inside a run-time for/while/if. That earlier decision was " - "made ONCE, from the original value {value}, and cannot re-run when `{var}` " - "changes -- the compiled kernel would silently keep the stale decision.", + PHASE_AUTO_PROMOTE_DISABLED = ( + "`{var}` (a {meta}, value {value}) is updated inside a for/while/if " + "whose path is decided at run time. Keeping it correct there requires " + "promoting it to a {staged}, and automatic Meta-to-Staged promotion " + "is disabled (`CUTE_DSL_AUTO_M2S` is off, the default). Without the " + "promotion the compiler traces the body once and later iterations " + "would silently observe a wrong value, so it refuses instead.", ( - "Make `{var}` a {staged} BEFORE the run-time for/while/if, e.g. " - "`{var} = {type}({var})`, so the decision at {fold_location} becomes a " - "run-time branch that sees the updated value.", - "If `{var}` is a fixed setting, do not modify it inside the run-time " - "for/while/if (keep the update outside, or guard it with " - "`const_expr(...)`).", + "Make `{var}` a {staged} before the for/while/if, e.g. " + "`{var} = {type}(...)`, so the update is tracked explicitly.", + "Or opt in to automatic promotion with `CUTE_DSL_AUTO_M2S=True`.", ), ) PHASE_PYTHON_THEN_TRACKED = ( @@ -1543,14 +1614,6 @@ class DiagId(_DiagMixin, enum.Enum): "result is consistently one type.", ), ) - SCOPE_READ_NEVER_SET = ( - "A variable is read on a path where it was never given a value. It must be " - "set before it is read, on every path that can reach this point.", - ( - "Set the variable before the for/while/if, or set it on every branch " - "that can reach this read.", - ), - ) SCOPE_DEL_LOOP_CARRIED = ( "`{var}` is removed with `del` inside a for/while/if, but it carries a " "value in from before the block. Deleting it drops that carry, so the " @@ -1560,14 +1623,201 @@ class DiagId(_DiagMixin, enum.Enum): "value instead, or move the `del` outside the for/while/if.", ), ) - UNSUP_WALRUS_TUPLE_REBIND = ( - "A `:=` inside this tuple assignment rewrites `{var}` while another " - "element of the same right-hand side reads it out of order. Under staged " - "tracing the walrus's rebind retroactively changes that read, so the " - "elements would disagree with plain Python.", + SCOPE_READ_NEVER_SET = ( + "A variable is read on a path where it was never given a value{detail}. " + "It must be set before it is read, on every path that can reach this " + "point.", ( - "Assign the walrus result on its own line before the tuple, e.g. " - "`{var} = ` and then reference `{var}` in the tuple.", + "Set the variable before the for/while/if, or set it on every branch " + "that can reach this read.", + ), + ) + SCOPE_UNBOUND_NAME_IN_TRACE = ( + "`{var}` is read here but was never given a value on the traced path.", + ( + "Set `{var}` before this read on every path that can reach it (a " + "branch selected off at trace time does not bind it).", + ), + ) + SCOPE_LOCALS_SYNTH_MISS = ( + "`locals()` inside a rewritten if/ifexp/while arm sees only the names " + "the arm rebinds, and `{name}` is not one of them -- the source " + "program's `locals()` would see the whole function scope, so this " + "lookup cannot be answered faithfully.", + ( + "Read `{name}` directly instead of through `locals()`, or call " + "`locals()` outside the for/while/if arm.", + ), + ) + MEMREF_INREGION_ALLOC_REBIND = ( + "`{name}` is re-pointed inside staged control flow ({region_kind}) " + "to a handle rooted at memory ALLOCATED inside the region, and the " + "escape cannot be proven faithful: the previous handle is still " + "held by another live binding, the handle is a raw pointer or view " + "(no declared memory-space fact), or the space is not per-thread " + "registers. In-region allocations are hoisted to ONE function-entry " + "buffer when lowering, so an escaped old handle would alias the " + "fresh buffer and observe its overwrites instead of the buffer " + "Python named.", + ( + "Drop the extra binding of the old buffer (finish reading it " + "before re-pointing `{name}`), carry the tensor itself instead " + "of a raw pointer or view into it, allocate the backing memory " + "once outside the dynamic region and write `{name}` in place, " + "or make the selection a Python-time (constexpr) branch.", + ), + ) + MEMREF_STALE_SCRATCH_CONSUMED = ( + "`{name}` holds iteration-private scratch memory (re-pointed inside " + "staged control flow) and a superseded handle to it escapes that " + "iteration's window at {consumer_loc}. All in-region allocations " + "alias ONE function-entry buffer when lowering, so the escaped " + "handle would observe later overwrites instead of the buffer Python " + "read.", + ( + "Consume the scratch inside the iteration that filled it " + "(before the next re-point), or allocate the backing memory " + "once outside the dynamic region and write `{name}` in place.", + ), + ) + BOUNDARY_META_LOOP_CARRY = ( + "`{name}` is a Python value changed inside this for/while body by " + "code the compiler cannot see (a plain, non-jit method or function). " + "Its per-iteration update ran on plain Python values, so the carried " + "update cannot be reconstructed -- iterations after the first would " + "silently observe a wrong value.", + ( + "Decorate the function that updates `{name}` with the DSL's jit " + "decorator so the update is tracked, or make `{name}` a Runtime " + "value (Staged value) before the loop, e.g. `Int32(...)`.", + ), + ) + BOUNDARY_FLIP_GUARD_STAGED = ( + "`{name}` is now being changed to a Runtime value, but it gates an " + "update of `{flipped}` performed by a plain, non-jit function " + "(defined at {def_file}:{def_line}) inside a for/while body. That " + "update was compiled as an unconditional per-iteration store on the " + "proof that `{name}` never changes during the loop; changing it " + "breaks that proof, so iterations could silently observe a wrong " + "`{flipped}`.", + ( + "Decorate the function that updates `{flipped}` with the DSL's " + "jit decorator so the update is tracked, or make both `{name}` " + "and `{flipped}` Runtime values (Staged values) before the loop.", + ), + ) + BOUNDARY_CLOSURE_READ_THEN_PROMOTED = ( + "`{name}` is read through a closure by a plain, non-jit function " + "(defined at {def_file}:{def_line}) inside staged control flow, and " + "is now being changed so it must become a Runtime value. That " + "closure read ran on the plain Python value and was baked into the " + "kernel at its trace-time value; no rewrite can retarget it to the " + "staged update, so every use would keep observing the old value -- " + "matching Python for some runtime inputs and silently diverging for " + "others.", + ( + "Decorate the function that reads `{name}` with the DSL's jit " + "decorator, pass `{name}` to it as an argument, or make it a " + "Runtime value (e.g. `{name} = Int32({name})`) before the first " + "call.", + ), + ) + BOUNDARY_CLOSURE_READ_THEN_WRITTEN = ( + "`{name}` is changed here inside staged control flow, but it was " + "already read through a closure by a plain, non-jit function " + "(defined at {def_file}:{def_line}). That closure read ran on the " + "plain Python value and was baked into the kernel at its trace-time " + "value; no rewrite can retarget it to the staged update, so every " + "use would keep observing the old value -- matching Python for some " + "runtime inputs and silently diverging for others.", + ( + "Decorate the function that reads `{name}` with the DSL's jit " + "decorator so the read is tracked, or make `{name}` a Runtime " + "value (Staged value) before the plain function reads it.", + ), + ) + BOUNDARY_SHORT_CIRCUIT_EFFECT = ( + "`{name}` is changed by a call inside a short-circuited `and`/`or` " + "operand. The trace evaluates that operand exactly once, so the " + "change cannot be conditioned on the runtime value of the guarding " + "expression -- it would apply unconditionally, silently.", + ( + "Move the call out of the `and`/`or` expression to its own " + "statement, or guard it with an explicit `if`.", + ), + ) + BOUNDARY_MEMOIZED_IN_STAGED_CF = ( + "`{name}` is a functools.lru_cache-wrapped function called inside " + "staged (runtime) control flow. A cache hit at trace time skips the " + "wrapped body, so its work cannot be replayed for each runtime " + "execution of this region -- executions after the first would " + "silently observe stale effects.", + ( + "Hoist the `{name}` call out of the staged for/while/if, or call " + "the undecorated function (`{name}.__wrapped__`) inside the " + "kernel.", + ), + ) + PHASE_META_FIELD_CHANGED_IN_CF = ( + "Meta-primitive field `{owner_class}.{attr}` cannot change across " + "iterations of staged control flow ({old_value!r} -> {new_value!r}). " + "The compiler traces the body only once and would silently discard " + "later changes.", + ( + "Use a DSL Numeric type for `{attr}` so it is tracked by " + "`pyir.ref`, or hoist the assignment outside staged control " + "flow.", + ), + ) + CONTAINER_META_REBUILT_IN_CF = ( + "The {kind} field `{var}` holds plain Python values (Meta values) " + "and is rebuilt inside staged control flow ({old_value!r} -> " + "{new_value!r}). The compiler traces the body only once, so the " + "rebuilt {kind} would silently keep its first-iteration values.", + ( + "Initialize the {kind} with DSL Numeric values (e.g. " + "`Int32(0)`) so each element is tracked. In a staged `for` " + "body, `CUTE_DSL_AUTO_M2S=True` promotes the elements " + "automatically.", + "If the {kind} is trace-time structure, rebuild it in a " + "`range_constexpr(...)` loop instead.", + ), + ) + CONTAINER_SET_REBUILT_IN_CF = ( + "The set field `{var}` is rebuilt inside staged control flow " + "({old_value!r} -> {new_value!r}). Set members are compile-time " + "structure (hashing and membership are decided while tracing), so " + "the rebuilt set would silently keep its first-iteration members.", + ( + "Store per-iteration runtime values in a tuple, list, or dict " + "field instead; a set cannot hold runtime values.", + "If the set is trace-time structure, rebuild it in a " + "`range_constexpr(...)` loop instead.", + ), + ) + UNSUP_META_CONTAINER_MUTATION = ( + "In-place mutation `{container}.{method}(...)` of a meta Python " + "{kind} is not allowed inside staged control flow. The compiler " + "traces the body once and would silently discard the per-iteration " + "mutation.", + ( + "Build the {kind} before the staged region, or use a DSL " + "collection/struct type that the compiler can track for " + "accumulation.", + ), + ) + SCOPE_READ_OF_SUPERSEDED_GENERATION = ( + "`{var}` is {access} through an object that was replaced by a whole-object " + "assignment{detail}. The compiler keeps ONE live storage cell per field " + "across such a replacement, so keeping BOTH the old and the new object " + "alive does not work: this access would silently observe the replacing " + "object's current value instead of the value the old object held when it " + "was captured.", + ( + "Copy the fields you need into plain locals BEFORE the whole-object " + "assignment, e.g. `saved = obj.field`.", + "Or access the field through the live binding instead of the retained " + "handle.", ), ) CONTAINER_TUPLE_LENGTH_CHANGED = ( @@ -1576,6 +1826,18 @@ class DiagId(_DiagMixin, enum.Enum): "followed.", ("Keep `{var}` the same length on every branch and every loop pass.",), ) + CONTAINER_TUPLE_SUBCLASS_NOT_REBUILDABLE = ( + "A `{type}` value carried through a for/while/if must be rebuilt from " + "its parts, but `{type}` cannot be reconstructed that way{detail}. " + "Rebuilding it as a plain tuple would silently drop its named fields / " + "methods, so this is refused.", + ( + "Use a `typing.NamedTuple` (or a plain tuple) for values carried " + "through a for/while/if.", + "Or give `{type}` a constructor that accepts the item iterable " + "(like `tuple` itself) and no extra per-instance state.", + ), + ) CONTAINER_STRUCTURE_CHANGED = ( "`{var}` has a different structure at the end of this `{op_type}` than at " "the start{detail}. A value carried through a `{op_type}` must keep the same " @@ -1602,39 +1864,145 @@ class DiagId(_DiagMixin, enum.Enum): "`{var}.field = new_value`.", ), ) - CONTAINER_LIST_SHAPE_MUTATED = ( - "`{var}` is a list whose length changes (e.g. `.append(...)`/`.pop()`) " - "inside a for/while/if whose path is decided at run time. The body is " - "traced once, so the per-iteration shape change would be silently lost. " - "Only a fixed-shape container can be carried through run-time control flow.", + CONTAINER_DICT_META_WRITE_UNPROMOTED = ( + "`{var}` is a dict entry holding a {meta} (value {old_value}) that is " + "changed (to {new_value}) inside a for/while/if whose path is decided at " + "run time, but no earlier use let the compiler promote it to a {staged}. " + "The body runs once at compile time, so the per-iteration change would be " + "silently discarded.", ( - "Build the list before the run-time for/while/if, or accumulate into a " - "{staged} container the compiler can track.", + "Store a {staged} in the entry from the start, e.g. " + "`{var} = Int32(0)`, so updates inside the for/while/if are tracked.", + "If the entry is a fixed compile-time setting, set it once before the " + "for/while/if.", ), ) CONTAINER_DICT_KEY_SET_MUTATED = ( - "`{var}` changes its set of keys inside a for/while/if whose path is " - "decided at run time (a key created/removed, or a mapping subclass like " - "`Counter`/`defaultdict` that materialises keys on access). The body is " - "traced once, so the per-iteration key-set change cannot be followed.", + "the key set of dict `{var}` is changed ({detail}) inside a for/while/if " + "whose path is decided at run time. The body runs once at compile time, so " + "a per-iteration key insertion/removal would be silently discarded.", ( - "Seed every key before the run-time for/while/if and only UPDATE " - "existing keys inside it; do not add/remove keys or use a " - "`__missing__`-backed mapping there.", + "Create every key before the for/while/if and update the VALUES inside it.", + "If the loop is a compile-time unroll, use `range_constexpr` so the " + "mutation is realized at trace time.", ), ) - CONTAINER_SUBSCRIPT_WRITE_UNTRACKED = ( - "An element of a `{py_type}` is written with `[...] = ...` inside a " - "for/while/if whose path is decided at run time, but a `{py_type}` is a " - "plain {meta}: the body is traced once, so the per-iteration write would " - "be silently lost. Only a dict, a list, or a {staged} container can be " - "updated by element there.", + CONTAINER_DICT_KEY_STAGED = ( + "a {staged} is used as a dict KEY on tracked dict `{var}`. Keys name " + "compile-time storage places, so they must be fixed Python values " + "(str/int/tuple); a runtime value cannot name a place.", + ( + "Use a compile-time key (a Python str/int), or restructure so the " + "runtime value is the ENTRY, not the key.", + ), + ) + CONTAINER_DICT_KEY_STAGED_HASH = ( + "the `__hash__` of a KEY on tracked dict `{var}` consumed a {staged}. " + "Keys name compile-time storage places, so a key whose hash is derived " + "from a runtime value cannot name a place.", + ( + "Hash only fixed Python values in the key's `__hash__`, or use a " + "compile-time key (a Python str/int) and keep the runtime value as " + "the ENTRY, not the key.", + ), + ) + CONTAINER_LIST_META_WRITE_UNPROMOTED = ( + "`{var}` is a list element holding a {meta} (value {old_value}) that is " + "changed (to {new_value}) inside a for/while/if whose path is decided at " + "run time, but no earlier use let the compiler promote it to a {staged}. " + "The body runs once at compile time, so the per-iteration change would be " + "silently discarded.", + ( + "Store a {staged} in the element from the start, e.g. " + "`{var} = Int32(0)`, so updates inside the for/while/if are tracked.", + "If the element is a fixed compile-time setting, set it once before " + "the for/while/if.", + ), + ) + CONTAINER_LIST_SHAPE_MUTATED = ( + "the length or element order of list `{var}` is changed ({detail}) inside " + "a for/while/if whose path is decided at run time. The body runs once at " + "compile time, so a per-iteration append/remove/reorder would be silently " + "discarded.", + ( + "Build the list before the for/while/if and update its ELEMENTS " + "inside it (`{var}[i] = ...`).", + "If the loop is a compile-time unroll, use `range_constexpr` so the " + "mutation is realized at trace time.", + ), + ) + CONTAINER_LIST_INDEX_STAGED = ( + "a {staged} is used as the INDEX into tracked list `{var}`. Indices name " + "compile-time storage places, so they must be fixed Python integers; a " + "runtime value cannot name a place.", ( - "Use a dict or list for run-time element updates, or a {staged} " - "container the compiler can track.", - "Or keep the for/while/if fully compile-time with " - "`range_constexpr(...)` / `const_expr(...)` so the container can stay " - "a {meta}.", + "Use a compile-time index (a Python int), or restructure so the " + "runtime value is the ELEMENT, not the index.", + ), + ) + CONTAINER_SUBSCRIPT_WRITE_UNTRACKED = ( + "`{var}` (a {kind}) has an element written inside a for/while/if whose " + "path is decided at run time. A {kind} holds compile-time Python storage " + "the compiler cannot track per iteration, so the body -- which runs once " + "at compile time -- would silently discard the per-iteration write.", + ( + "Use a plain Python `list` (tracked per element) or a {staged} " + "value instead of the {kind}.", + "If the loop is a compile-time unroll, use `range_constexpr` so the " + "write is realized at trace time.", + ), + ) + CONTAINER_DICT_KEY_SET_BAKED_READ = ( + "membership/lookup over dict `{var}` is consumed at compile time, but " + "its key set was changed (created key(s) {detail}) inside a " + "for/while/if whose path is decided at run time. The body ran once at " + "compile time, so the consumed key set reflects that one pass -- the " + "kernel would silently behave as if the branch was taken, whatever " + "the runtime values.", + ( + "Create every key before the for/while/if and update the VALUES inside it.", + "If the for/while/if is compile-time, use `range_constexpr(...)` " + "/ `const_expr(...)` so the insertion is realized at trace time.", + ), + ) + CONTAINER_OPAQUE_SUBSCRIPT_KEY_CREATED = ( + "an element write through `{var}` (a {kind}, whose `__setitem__` the " + "compiler cannot see into) created container entry/entries {detail} " + "inside a for/while/if whose path is decided at run time. The body " + "runs once at compile time and a created entry has no storage cell " + "from before the for/while/if, so it cannot follow the runtime path " + "-- later reads would silently observe it regardless of the branch " + "taken.", + ( + "Create the entry before the for/while/if and update its VALUE inside it.", + "Use a plain Python `dict`/`list` (tracked per element) instead " + "of the {kind}.", + "If the for/while/if is compile-time, use `range_constexpr(...)` " + "/ `const_expr(...)` so the write is realized at trace time.", + ), + ) + CONTAINER_DICT_GET_MISS_IN_STAGED_CF = ( + "`.get()` on tracked dict `{var}` missed key {detail} inside a " + "for/while/if whose path is decided at run time, while other entries of " + "the dict are updated there. The miss default would bake as a fixed " + "compile-time value on this path.", + ( + "Create the key before the for/while/if so every path reads a tracked " + "entry.", + ), + ) + CONTAINER_DICT_ITERATED_IN_STAGED_CF = ( + "tracked dict `{var}` is iterated ({method}) inside a for/while/if " + "whose path is decided at run time, after its key set was changed " + "there (created key(s) {detail}). The body ran once at compile time, " + "so the enumerated key set reflects that one pass -- the kernel " + "would silently iterate entries as if the path creating them was " + "taken, whatever the runtime values.", + ( + "Create every key before the for/while/if and update the VALUES " + "inside it -- iteration over a stable key set is supported.", + "If the for/while/if is compile-time, use `range_constexpr(...)` " + "/ `const_expr(...)` so the insertion is realized at trace time.", ), ) PHASE_CONVERSION_FAILED = ( @@ -1705,6 +2073,69 @@ class DiagId(_DiagMixin, enum.Enum): "function.", ("Pass `{name}` in as a function argument instead.",), ) + UNSUP_YIELD = ( + "`yield` makes a function a generator, which cannot be compiled: calling " + "it only creates a generator object and never executes the body.", + ( + "Define the generator at module scope (plain Python, outside compiled " + "code) and consume it there, or build a list instead.", + ), + ) + UNSUP_ASYNC = ( + "`async`/`await` constructs cannot be compiled: a coroutine body does not " + "execute when called and there is no event loop to drive it here.", + ("Compute the value with a plain (non-async) call instead.",), + ) + UNSUP_EXCEPT_STAR = ( + "`except*` (ExceptionGroup handling) is not supported in a compiled function.", + ("Use a plain `except` clause instead.",), + ) + UNSUP_DEL_IN_STAGED_CF = ( + "`{obj}.{attr}` is deleted inside a for/while/if that runs on the GPU, but " + "the deletion happens once while building the program -- it cannot be made " + "conditional on the region's runtime condition.", + ( + "Move the `del` (or `delattr`) outside the for/while/if.", + "If the attribute must differ per path, assign it a new value instead " + "of deleting it.", + ), + ) + PROPERTY_GETTER_MUTATES_IN_STAGED_CF = ( + "Reading `{obj}.{attr}` runs the property getter `{getter}`, which " + "writes state, and the read is inside a for/while/if that runs on the " + "GPU: the getter executes once while building the program, so its side " + "effect cannot re-run per iteration and both the result and the " + "mutated state would freeze at their first values.", + ( + "Make the getter pure and perform the update in an explicit " + "method call whose effect assigns to tracked state.", + "If the side effect is trace-time-only bookkeeping, read the " + "backing field directly instead of going through the property.", + ), + ) + UNSUP_GLOBAL_WRITE_IN_STAGED_CF = ( + "`{callee}` writes the module global `{name}`, and it is called inside a " + "for/while/if that runs on the GPU: the write executes once while " + "building the program, not once per iteration, so every later read of " + "`{name}` would see a value frozen at its first update.", + ( + "Pass the state in as an argument and return the updated value " + "instead of mutating a module global.", + "If the call is trace-time configuration whose result never feeds " + "computed values, move it outside the for/while/if.", + ), + ) + UNSUP_IMPORT_IN_STAGED_CF = ( + "`{stmt}` is inside a for/while/if that runs on the GPU, and the module is " + "not imported yet: its code would execute once while building the program, " + "not under the region's runtime condition.", + ( + "Import the module at module scope (or before the for/while/if); an " + "already-imported module is a cache lookup and is allowed here.", + "If a per-path module choice is needed, import every candidate outside " + "the for/while/if and select between them.", + ), + ) UNSUP_MIXED_ASSIGN_TARGETS = ( "An assignment that mixes plain names, subscripts, and tuple targets on a " "single line is not supported here.", @@ -1716,6 +2147,27 @@ class DiagId(_DiagMixin, enum.Enum): ("Save the function to a .py file and import it from there.",), ) + UNSUP_WALRUS_CONDITIONAL = ( + "A walrus assignment (`{name} := ...`) inside a conditionally-evaluated " + "expression (an `and`/`or` right-hand side, a ternary branch, or a " + "comprehension) is not supported in compiled code: whether the write " + "happens cannot be decided at compile time.", + ( + "Move the walrus assignment out to its own statement before the " + "expression, or compute the value into a variable first.", + ), + ) + UNSUP_WALRUS_EVAL_ORDER = ( + "A call (`{call}`) that Python evaluates around the walrus assignment " + "(`{name} := ...`) in this statement cannot be kept in its original " + "evaluation order by the compiler (it sits in a conditionally-evaluated " + "position, or between two walrus assignments).", + ( + "Move the walrus assignment out to its own statement before the " + "expression, or compute the call's result into a variable first.", + ), + ) + # ===================================================================== # Migrated from raw DSLRuntimeError raises across the DSL (author # mistakes: config, types, arguments, calls, launch, tensors, ...). @@ -1753,6 +2205,15 @@ class DiagId(_DiagMixin, enum.Enum): "Passed {got} argument(s) to FFI function, but it expects {expected}", ("Check the function signature and provide the correct number of arguments",), ) + ARG_CONSTEXPR_MISMATCH = ( + "Argument `{arg_name}` was `None` when this function was compiled, and it was " + "baked into the compiled function as constexpr. Expected exactly the same " + "argument `None`, but got `{arg_type}`.", + ( + "Pass `None` for `{arg_name}`, as at compile time.", + "Or compile the function again with the value you want to pass.", + ), + ) ARG_FOR_LOOP_STEP_NOT_INT = ( "Loop bounds and step must be integers, not {kind}", ( @@ -1843,6 +2304,26 @@ class DiagId(_DiagMixin, enum.Enum): "Convert the value to an integer type, e.g., `int({name})`.", ), ) + # --- ATTR --- + ATTR_BUILDER_REQUIRES_CONSTANT = ( + "`{var}` is passed to `{callee}`, which needs one plain compile-time " + "value, but `{var}` changes inside a for/while/if -- no single value " + "exists to pass.", + ( + "Pass a {staged} value through a DSL API that accepts runtime " + "values, or keep `{var}` a fixed compile-time constant.", + ), + ) + # --- BODY --- + BODY_BORN_SEED_UNPLACEABLE = ( + "`{var}` is first given a value inside a for/while/if, but that value " + "was computed on a different path, so no store can be placed where the " + "assignment actually runs.", + ( + "Assign `{var}` a value computed on the same path (or before the " + "for/while/if) so the assignment can be compiled where it runs.", + ), + ) # --- CALL --- CALL_BUILTIN_KWARGS_UNSUPPORTED = ( "The built-in function '{fcn}' does not support keyword arguments.", @@ -1876,6 +2357,11 @@ class DiagId(_DiagMixin, enum.Enum): "No function was provided to compile. Pass a callable decorated with @cute.jit.", ("Ensure you pass a valid @cute.jit-decorated function to cute.compile().",), ) + CALL_KERNEL_TARGET = ( + "`{function_name}` is decorated with @cute.kernel. cute.compile() " + "expects a @cute.jit-decorated function to compile its target as the host entry point.", + ("Compile the @cute.jit function that launches `{function_name}`.",), + ) CALL_MISSING_ARG = ( "Required argument `{name}` is missing in the call to `{function_name}`.", ("Pass a value for `{name}`.",), @@ -1935,10 +2421,11 @@ class DiagId(_DiagMixin, enum.Enum): ), ) CONFIG_ATTRIBUTES_UNSUPPORTED = ( - "The `@kernel` decorator with `attributes=` is not supported by this DSL. Only experimental CuTe supports this feature.", + "Non-empty `@kernel` attributes require CuTe extension compilation.", ( + "Make `attributes=` resolve to `None` or an empty dict for this kernel.", "Remove the `attributes=` parameter from the @kernel decorator.", - "Or use `@cute.experimental.kernel(attributes=...)` if you need kernel-level attributes.", + "Or enable CuTe extension compilation for this program.", ), ) CONFIG_ATTR_KEY_UNSUPPORTED = ( @@ -2007,9 +2494,9 @@ class DiagId(_DiagMixin, enum.Enum): ("Use `Constexpr` annotation or a Python constant for max_number_threads.",), ) CONFIG_MISSING_NVDISASM = ( - "{vars} requires the 'nvdisasm' tool to write SASS output, but it is not found in PATH.", + "{vars} requires the 'nvdisasm' tool to write SASS output, but no compatible binary was found.", ( - "Install the CUDA Toolkit from https://developer.nvidia.com/cuda-downloads; if CUDA is installed, add its bin directory to PATH: export PATH=/usr/local/cuda/bin:$PATH", + "Install it with `pip install nvidia-cutlass-dsl[sass]`, or install a CUDA Toolkit and expose it via CUDA_HOME/CUDA_PATH.", ), ) CONFIG_MISSING_TVM_FFI = ( @@ -2068,6 +2555,19 @@ class DiagId(_DiagMixin, enum.Enum): 'Pass a version string like "12.3" or use DSLCudaVersion("{version_string}")', ), ) + # --- INSTRUMENTATION --- + INSTRUMENTATION_GAP = ( + "Function `{name}` is being traced as a compilation entry point, but " + "it was never rewritten by the DSL preprocessor and is not declared " + "native -- its variable reads/writes would be invisible to the " + "compiler, producing silently wrong code.", + ( + "Decorate `{name}` with @jit / @kernel so the preprocessor " + "instruments it before tracing.", + "Or opt it out explicitly with `preprocess=False` if it must run " + "as plain Python.", + ), + ) # --- LAUNCH --- LAUNCH_INVALID_CLUSTER = ( "Launch cluster must have exactly 3 dimensions.", @@ -2104,7 +2604,84 @@ class DiagId(_DiagMixin, enum.Enum): "If you meant to discard the call, remove it.", ), ) + # --- OWNER --- + OWNER_CLASS_CHANGED = ( + "An object whose state is tracked across a for/while/if changed its " + "class from `{old_class}` to `{new_class}` -- the tracked state was " + "recorded under `{old_class}` and cannot be re-derived for " + "`{new_class}`.", + ( + "Create a new `{new_class}` object instead of reassigning " + "`__class__` on one the compiler is tracking.", + ), + ) + OWNER_DEL_FINALIZER_AT_BIRTH = ( + "`{cls}` defines `__del__` (via `{definer}`): the finalizer runs at " + "a garbage-collection-determined instant, which has no position in " + "the traced program, so an object of this class created here cannot " + "be tracked.", + ( + "Release resources with an explicit close()/context manager " + "instead of `__del__` on `{definer}`.", + "Or create the object outside the traced function and pass it in.", + ), + ) + OWNER_DECLARED_SURFACE_VIOLATION = ( + "`{obj}` opted into the DSL value protocol " + "(`__extract_mlir_values__`), declaring which of its values follow " + "run-time control flow, but {violation}. State outside the " + "declaration is fixed once code is read and cannot change inside a " + "for/while/if whose path is decided at run time.", + ( + "Declare the state: make `__extract_mlir_values__` / " + "`__new_from_mlir_values__` carry it (and keep the DSL value " + "type's own `__hash__`).", + "Or perform the change outside the run-time for/while/if.", + ), + ) + OWNER_FABRICATED_ATTR_IN_STAGED_CF = ( + "`{obj}.{attr}` is fabricated by `{definer}.__getattr__` inside " + "staged (runtime) control flow: the name resolves to no storage " + "slot, so the trace-time value would bake here and silently mask " + "state that can change across runtime executions of this region.", + ( + "Read `{obj}.{attr}` into a local before the staged " + "for/while/if and use the local, or give `{attr}` real storage " + "(assign it in `__init__`).", + ), + ) # --- PHASE --- + PHASE_IDENTITY_ON_TRACKED = ( + "`is` asks for Python object identity, but `{what}` is a " + "compiler-tracked value here: the compiler re-wraps tracked values " + "while staging code, so wrapper identity does not follow the " + "program's own object identity.", + ( + "Compare values with `==` instead of `is`.", + "Identity tests against `None` (or another non-numeric sentinel) " + "stay exact and are supported.", + ), + ) + PHASE_SERIALIZE_STAGED = ( + "`{what}` holds a staged (runtime) value: serializing it would bake " + "compiler trace state into bytes that cannot round-trip to the value " + "computed at run time.", + ( + "Serialize plain Python values computed outside staged code.", + "If the value is needed at run time, keep it in scalars/tensors " + "instead of serialized bytes.", + ), + ) + PHASE_NUMERIC_PROTOCOL_ON_STAGED = ( + "`{proto}` is not supported on `{what}`, a {staged}: this numeric " + "protocol has no runtime (staged) form.", + ( + "For 3-argument pow compute `(a ** b) % m` when wraparound " + "semantics are acceptable.", + "For `@` and complex(), keep the operands plain Python values -- " + "Python defines no scalar `@` either.", + ), + ) PHASE_CONDITIONAL_NOT_DYNAMIC = ( "The condition must be a {staged} value, not a {meta}.", ("Ensure the condition is a runtime value (Boolean, Int32, etc.).",), @@ -2130,6 +2707,68 @@ class DiagId(_DiagMixin, enum.Enum): "If the value is only known at run time, use the runtime form instead (for example a runtime `if`/loop or a runtime assert).", ), ) + PHASE_STRUCTURAL_CONSTANT_MUTATED = ( + "`{var}` is changed here so it must become a {staged}, but it was " + "already consumed as a compile-time constant (value {value}, at " + "{read_file}:{read_line}) inside a for/while/if -- as a loop trip " + "count, a Python-sequence index, a folded if/while predicate, or a " + "similar structural use. That structure is fixed at compile time, " + "so every iteration would keep the old value silently.", + ( + "Make `{var}` a {staged} from the start (e.g. `Int32(...)`) and " + "use a runtime loop/index, or keep it a fixed compile-time " + "constant and do not change it inside the for/while/if.", + ), + ) + PHASE_ARM_LOCAL_CONSTANT_ESCAPES = ( + "`{var}` is read here, but it was changed as a compile-time constant " + "inside an if branch that only runs on some GPU paths (at " + "{write_file}:{write_line}). The compiler runs the body once, so " + "this read would keep that branch's value even when the branch is " + "skipped at run time.", + ( + "Make `{var}` a {staged} before the if (e.g. `{var} = " + "Int32(...)`) so the update is carried on every path.", + "Or keep every read of `{var}` inside the same if branch that " + "changes it.", + ), + ) + PHASE_PREDICATE_FOLDED_STALE = ( + "`{var}` is changed here inside a for/while body, but an if/while " + "condition already folded on its compile-time value ({value}, at " + "{read_file}:{read_line}); the folded branch was fixed at trace " + "time, so every iteration would keep the stale decision silently.", + ( + "Make the value a {staged} at construction (e.g. `Int32(...)`) so " + "the condition stays a runtime comparison.", + "Or hoist the mutation out of the loop and keep the value a fixed " + "compile-time constant.", + ), + ) + # --- READ --- + READ_DEPTH_OVERFLOW = ( + "The attribute/subscript chain `{var}` is {depth} accesses deep, " + "which exceeds the {budget} accesses the compiler tracks per " + "statement -- deeper reads cannot be attributed to their storage " + "and would bake silently.", + ( + "Split the chain across statements with intermediate variables " + "so each statement stays within {budget} accesses.", + ), + ) + # --- RECONSTRUCT --- + RECONSTRUCT_NO_PROTOCOL = ( + "A value of type `{cls}` must be rebuilt here from its compiled " + "value, but `{cls}` provides no way to do so: it implements neither " + "`__new_from_mlir_values__` nor a constructor accepting the compiled " + "value.", + ( + "Implement `__extract_mlir_values__` / `__new_from_mlir_values__` " + "on `{cls}` so the compiler can rebuild it.", + "Or carry a DSL value type (Int32, Float32, TensorSSA, ...) " + "across the for/while/if instead.", + ), + ) # --- SCOPE --- SCOPE_CLOSURE_CAPTURE = ( "Function `{func_name}` captures variable `{var_name}`, which is not supported in staged for/while/if.", @@ -2138,28 +2777,63 @@ class DiagId(_DiagMixin, enum.Enum): "Define the function inside the loop/if, or refactor to avoid the closure.", ), ) - SCOPE_LAMBDA_CAPTURE = ( - "A `lambda` invoked inside a staged for/while/if reads captured variable " - "`{var_name}` from the enclosing function. The lambda body runs as plain " - "Python while your code is read (at trace time), so `{var_name}` is frozen " - "at its first-pass value instead of following updates made inside the " - "staged region.", + SCOPE_NONLOCAL_WRITE_IN_STAGED_CF = ( + "`{name}` was already updated through a `nonlocal` write in another " + "function scope inside this dynamic for/while/if, and is accessed " + "here from a different scope of the same region. The two accesses " + "cannot be routed to one storage cell across the region, so one " + "side's update would be silently lost at the join.", ( - "Wrap the lambda with `@cute.jit`, e.g. `f = cute.jit(lambda ...: ...)`, " - "so its body is traced on every pass.", - "Or pass `{var_name}` in as a lambda argument instead of capturing it.", + "Return the new value from the nested function and rebind it at " + "the call site (`{name} = fn(...)`), or move the `nonlocal` " + "write outside the dynamic for/while/if.", ), ) - SCOPE_CTXMGR_TRACE_ONLY = ( - "The `with` uses context manager `{ctx_type}`, whose `__enter__`/`__exit__` " - "are plain Python and run only once, while your code is read (at trace " - "time). Inside a staged for/while/if the block runs on every pass at run " - "time, so those effects would be frozen at the first pass and the compiled " - "code would not match Python.", + # --- SNAPSHOT --- + SNAPSHOT_UNMATERIALIZABLE = ( + "`{var}` holds a value captured earlier in the trace, but the tracked " + "variable it was captured from has moved on since, and the captured " + "value itself was computed inside a for/while/if region that has " + "already closed -- so neither the captured value nor a faithful " + "re-read of it is available here.", ( - "Decorate both `{ctx_type}.__enter__` and `{ctx_type}.__exit__` with " - "`@cute.jit` so they are traced on every pass.", - "Or move the `with` out of the staged for/while/if.", + "Read the source variable directly at this point instead of " + "keeping a captured copy across the for/while/if.", + "Or capture the value before the for/while/if so the copy stays " + "available on every path.", + ), + ) + # --- SPEC --- + SPEC_DESCRIPTOR_HOP = ( + "Re-validating this compiled function would execute the `{attr}` " + "descriptor of `{owner_type}` while re-resolving `{path}`; a " + "specialization row must re-resolve through plain storage only. " + "This row should have been classified trace-internal at record time " + "-- a bug in the DSL, not a mistake in your code.", + ( + "Recompile via cute.compile (or re-invoke the @jit function) so " + "the value is re-read on a fresh trace.", + "And report this diagnostic to the DSL team.", + ), + ) + SPEC_UNVERIFIABLE_REENTRY = ( + "This compiled function baked a trace-time value that cannot be " + "re-checked against live state, so re-invoking the compiled handle " + "cannot be validated for reuse.", + ( + "Re-invoke the @jit-decorated function instead of the compiled " + "handle so the value is re-read on a fresh trace.", + "Or recompile with cute.compile after changing host state.", + ), + ) + # --- STALE --- + STALE_SPECIALIZATION = ( + "This compiled function was specialized on `{path}` = `{baked}`, but " + "the value at this call is `{live}` -- the compiled code still uses " + "the old value.", + ( + "Recompile via cute.compile (or re-invoke the @jit function) " + "after changing `{path}`.", ), ) # --- TENSOR --- @@ -2181,6 +2855,18 @@ class DiagId(_DiagMixin, enum.Enum): ), ) # --- TYPE --- + TYPE_CHANGED_INSIDE_REGION = ( + "`{var}` changes its compiled type from `{old_type}` to `{new_type}` " + "inside {region}. A value carried across a for/while/if boundary must " + "keep one compiled type, so this change cannot be carried out of the " + "region.", + ( + "Hoist the type change out of the for/while/if so `{var}` has the " + "new type on every path.", + "Or make the types agree: convert explicitly so every assignment " + "to `{var}` produces the same compiled type.", + ), + ) TYPE_CLUSTER_NOT_INT = ( "The {config_name} dimensions must be integers.", ("Provide integer values for {config_name}.",), @@ -2266,6 +2952,19 @@ class DiagId(_DiagMixin, enum.Enum): "The while loop inputs must be convertible to runtime values.", ("Ensure all inputs are DSL types or implement DynamicExpression.",), ) + # --- UNOBSERVED --- + UNOBSERVED_WRITE_POSITION_UNKNOWN = ( + "`{var}` was rebound by code the compiler cannot observe (no tracked " + "assignment recorded the new value), and a for/while/if boundary has " + "passed since the variable's last tracked access -- the write cannot " + "be placed at its true position in the compiled program.", + ( + "Perform the rebind with a plain assignment in jit-decorated code " + "so the write is tracked at its position.", + "Or move the rebind so no for/while/if boundary separates it from " + "the next use of `{var}`.", + ), + ) # --- UNSUP --- UNSUP_ARCH = ( "The `vectorize` attribute requires compute capability 10.0 or higher; your target is `{arch}`.", @@ -2304,6 +3003,30 @@ class DiagId(_DiagMixin, enum.Enum): "The `in` operator is not supported between these values.", ("Use a supported comparison, or restructure the check to avoid `in`.",), ) + # --- WRAPPER --- + WRAPPER_CLASS_MERGE = ( + "`{var}` leaves this for/while/if holding a `{new_cls}`, but entered " + "it holding a `{old_cls}` over the same compiled type. The two Python " + "classes have no common replacement, so reads after the region cannot " + "reconstruct one consistent value.", + ( + "Use one wrapper class for `{var}` on every path through the for/while/if.", + "Or convert explicitly before the for/while/if closes so both " + "paths produce the same class.", + ), + ) + WRAPPER_REBIND_DUPLICATE_LEAF = ( + "A `{cls}` is replaced inside {region}, but the value it held on " + "entry stores the SAME compiled value at more than one leaf " + "position. Carrying the replacement pairs old and new leaves by " + "value identity, so the repeated positions cannot be told apart and " + "an update could be silently wired to the wrong one.", + ( + "Enter the for/while/if with a distinct value at every leaf " + "position (compute each initial value separately), or update " + "the positions in place instead of replacing the whole object.", + ), + ) class WarnId(_DiagMixin, enum.Enum): @@ -2327,6 +3050,49 @@ class WarnId(_DiagMixin, enum.Enum): ), ) + TYPE_INT_LITERAL_OUT_OF_RANGE = ( + "The Python integer {value} does not fit in {type} (range " + "[{min}, {max}]); it was silently truncated to {wrapped}. This drops " + "the high bits of the original value.", + ( + "Use a wider integer type that can hold {value}, e.g. " + "`Int64({value})` or `Uint64({value})`.", + "If the wrap-around is intentional (e.g. materializing a specific " + "bit pattern), mask the value to the type width first, e.g. " + "`{type}({value} & 0x{mask:X})`, to make the intent explicit.", + ), + ) + + TYPE_FLOAT_TO_INT_OUT_OF_RANGE = ( + "The Python float {value} does not fit in {type} (range " + "[{min}, {max}]); it was silently narrowed to {result}. " + "The magnitude of the original value is lost.", + ( + "Clamp the value to the target range before converting, e.g. " + "`{type}(max({min}, min({max}, {value})))`.", + "Or use a wider integer type that can hold {value}.", + "Do not depend on the exact value {result}: narrowing a float the " + "target cannot represent has no portable definition.", + ), + ) + + TYPE_FLOAT_LITERAL_OVERFLOW = ( + "The Python float {value} is larger than the maximum finite {type} " + "value ({max:g}); it silently became {wrapped}. The magnitude of the " + "original value is lost.", + ("Use a wider float type that can hold {value}, e.g. `Float64({value})`.",), + ) + + TYPE_FLOAT_LITERAL_UNDERFLOW = ( + "The Python float {value} is smaller than the smallest nonzero {type} " + "value ({tiny:g}); it silently became {wrapped}. The original nonzero " + "value is lost.", + ( + "Use a wider float type that can represent {value}, e.g. " + "`Float64({value})`.", + ), + ) + def report_warning( warn_id: "WarnId", diff --git a/python/CuTeDSL/cutlass/base_dsl/dsl.py b/python/CuTeDSL/cutlass/base_dsl/dsl.py index d46e7e90de..b23b108dda 100644 --- a/python/CuTeDSL/cutlass/base_dsl/dsl.py +++ b/python/CuTeDSL/cutlass/base_dsl/dsl.py @@ -67,7 +67,12 @@ KeepCUBIN, ) from .ast_helpers import DSLOptimizationWarning -from .common import DSLRuntimeError, active_env_manager, target_version +from .common import ( + DSLRuntimeError, + DSLUserCodeError, + active_env_manager, + target_version, +) from .compile_backend import CompileContext, get_compiler_backend # ============================================================================= @@ -86,6 +91,7 @@ ) from .ast_preprocessor import DSLPreprocessor +from .pyir_class_facts import configure_facts_cache, record_jit_decoration from .pyir_preprocessor import PyIRDSLPreprocessor from .preprocess_mode import _PreprocessModeState from .common import * @@ -137,6 +143,12 @@ MLIR_DYNAMIC = -9223372036854775808 +# Optional PyIR runtime hook bound at module load; absent pyir/ub dialects degrade it to None. +try: + from .pyir_runtime import _verify_no_used_poison as _PYIR_VERIFY_POISON +except ImportError: + _PYIR_VERIFY_POISON = None # type: ignore[assignment] # optional-import degrade + # Keyword parameter a compiler provider declares to receive the compiled module # back. Keep this in sync with the keyword-only parameter named return_module on # BaseDSL.compile_and_jit and Compiler.compile_and_jit. @@ -349,6 +361,14 @@ def __init__( self.jit_cache: JitCacheDict = JitCacheDict( max_elems=self.envar.jit_cache_max_elems ) + # On-disk cache for the PyIR per-class write-fact parses (source + # content-hash keyed); follows this DSL's file-caching switch. + configure_facts_cache( + os.path.join( + get_default_generated_ir_path(self.envar.prefix), "pyir_class_facts" + ), + enabled=not self.envar.disable_file_caching, + ) self.host_jit_decorator_name: str = f"@{BaseDSL.jit.__name__}" self.device_jit_decorator_name: str = f"@{BaseDSL.kernel.__name__}" @@ -403,8 +423,6 @@ def __init__( self.preprocessor: DSLPreprocessor = preprocessor self._preprocess_mode = _PreprocessModeState(self) - if preprocess: - self._preprocess_mode.stamp(self.preprocessor) log().info(f"Initializing {name} DSL") log().debug(f"Logger initialized for {self.name}") @@ -469,12 +487,35 @@ def _get_dsl(cls) -> Self: @classmethod @contextmanager def enable_pyir(cls) -> Generator[None, Any, None]: - dsl = cls._get_dsl() - dsl._preprocess_mode.push_pyir() + """Select the PyIR frontend for every trace started inside the block. + + Frontend choice is a property of the *trace*, not of one DSL object. + A single module routinely mixes DSLs -- a ``@cute.experimental.jit`` + entry inlines ``@cute.jit`` library helpers -- and each DSL singleton + owns its own ``preprocessor`` and ``envar``. Arming only the singleton + that ``_get_dsl()`` happens to return leaves the other one tracing v1 + while PyIR-rewritten helpers emit ``pyir.*`` ops into its module, and + its pipeline then omits ``enable-pyir=true`` (both on + ``lir-to-cute-dsl`` and in the compile options), so the prelower never + runs and ``pyir.ref`` reaches LLVM translation. + """ + # Materialize the singleton for `cls` (and assert one exists when + # called on BaseDSL itself) before snapshotting the registry. + cls._get_dsl() + # Snapshot: a DSL built with preprocess=False owns no preprocessor and + # cannot switch frontends. A singleton created after this point keeps + # its own mode; both cute DSLs already exist once `cutlass` is imported. + pushed: list[Any] = [] try: + for dsl in list(DSLSingletonMeta._instances.values()): + if not hasattr(dsl, "preprocessor"): + continue + dsl._preprocess_mode.push_pyir() + pushed.append(dsl) yield finally: - dsl._preprocess_mode.pop() + for dsl in reversed(pushed): + dsl._preprocess_mode.pop() @staticmethod def _can_preprocess(**dkwargs: Any) -> bool: @@ -537,12 +578,16 @@ def _preprocess_and_replace_code(func: Any) -> None: def jit_runner( cls: type["BaseDSL"], executor_name: str, - frame: Any, + location: DSLLocation, *dargs: Any, **dkwargs: Any, ) -> Any: """ Decorator to mark a function for JIT compilation. + + ``location`` is the user's call site, already resolved to a value by + the caller via :meth:`get_location_from_frame`. + """ log().info("jit_runner") @@ -550,8 +595,19 @@ def jit_runner_decorator(func: Any) -> Any: # Run preprocessor that alters AST preprocess_enabled = BaseDSL._can_preprocess(**dkwargs) func._dsl_cls = cls - func._decorator_location = BaseDSL.get_location_from_frame(frame) + func._decorator_kind = ( + "kernel" if executor_name == "_kernel_helper" else "jit" + ) + func._decorator_location = location func._preprocess_enabled = preprocess_enabled + # F-COVER decoration-time attestation: an explicit preprocess + # opt-out declares the function native; otherwise the rewrite + # point stamps ``__pyir_rewritten__`` before tracing. + if not preprocess_enabled: + func.__pyir_native__ = True + # Decoration-time producer for the per-class self-field write-fact + # registry (O(1) until PyIR activates; see pyir_class_facts). + record_jit_decoration(func) if not hasattr(func, "_preprocessed") and not preprocess_enabled: func._preprocessed = True @@ -634,20 +690,30 @@ def jit(cls, *dargs: Any, **dkwargs: Any) -> Any: """ Decorator to mark a function for JIT compilation for Host code. """ - cur_frame = inspect.currentframe() - assert cur_frame is not None - frame = cur_frame.f_back - return BaseDSL.jit_runner(cls, "_func", frame, *dargs, **dkwargs) + return BaseDSL.jit_runner( + cls, + "_func", + BaseDSL.get_location_from_frame( + inspect.currentframe().f_back # type: ignore[union-attr] + ), + *dargs, + **dkwargs, + ) @classmethod def kernel(cls, *dargs: Any, **dkwargs: Any) -> Any: """ Decorator to mark a function for JIT compilation for GPU. """ - cur_frame = inspect.currentframe() - assert cur_frame is not None - frame = cur_frame.f_back - return BaseDSL.jit_runner(cls, "_kernel_helper", frame, *dargs, **dkwargs) + return BaseDSL.jit_runner( + cls, + "_kernel_helper", + BaseDSL.get_location_from_frame( + inspect.currentframe().f_back # type: ignore[union-attr] + ), + *dargs, + **dkwargs, + ) @abstractmethod def _kernel_helper(self, func: Any, *args: Any, **kwargs: Any) -> Any: @@ -1555,9 +1621,7 @@ def get_version(self) -> "hashlib._Hash": def get_module_hash(self, module: ir.Module, function_name: str) -> str: s = io.BytesIO() module.operation.write_bytecode(s) - for attr, value in self.envar.__dict__.items(): - if value is not None: - s.write(str(value).encode()) + s.write(self.envar.cache_key_str().encode()) # Add compile options to the hash s.write(self.compile_options.to_str().encode()) hash_obj = self.get_version().copy() @@ -1662,6 +1726,15 @@ def _run_trace_finalize_hooks(self, module: ir.Module, function_name: str) -> No f"Trace finalize hook failed: {hook_name}", cause=e ) from e + def _inspection_pipeline(self, passes: str) -> str: + """Frame ``passes`` for IR-inspection dumps: under PyIR, prepend the + ``pyir-prelower`` unless ``passes`` already starts with it, matching production.""" + if self.envar.enable_pyir and not passes.lstrip().startswith( + ("pyir-prelower", "convert-pyir-to-scf") + ): + passes = f"pyir-prelower,{passes}" + return f"builtin.module({passes})" + def _compile_clone_and_save( self, module: ir.Module, pipeline: str, label: str ) -> Any: @@ -1713,12 +1786,8 @@ def build_module(self, module: ir.Module, function_name: str) -> ir.Module: # Poison-read check: any used ub.poison value is a definite bug # (a value was read in code that does not dominate the write). - try: - from .pyir_runtime import _verify_no_used_poison - - _verify_no_used_poison(module) - except ImportError: - pass + if _PYIR_VERIFY_POISON is not None: + _PYIR_VERIFY_POISON(module) # Verify the module try: @@ -1792,6 +1861,22 @@ def build_ir_module() -> tuple[ir.Module, Any]: _jit_scope(), self._track_deferred_kernel_launches(), ): + # PyIR entry attestation (V-9), then trace-arg + # intake: host trace args become candidate holders + # and F-SPEC roots. + from .pyir_runtime import ( + _pyir_assert_entry_attested, + ) + + _pyir_assert_entry_attested(funcBody, self) + try: + from .pyir_runtime import _pyir_register_trace_args + + _pyir_register_trace_args( + ir_args, ir_kwargs, sig=sig, entry_func=funcBody + ) + except Exception: + pass result = funcBody(*ir_args, **ir_kwargs) default_ret_values = self.generate_default_return_values( ir.InsertionPoint.current @@ -1807,17 +1892,29 @@ def build_ir_module() -> tuple[ir.Module, Any]: err_filename = tb.tb_frame.f_code.co_filename err_lineno = tb.tb_lineno tb = tb.tb_next + # The curated funnel is a PyIR-mode fact: only the + # PyIR rewrite guarantees a genuinely-unbound read + # never escapes as a bare NameError. + if not self.envar.enable_pyir: + raise DSLUserCodeError( + f"NameError in `{funcBody.__name__}`: {name_error}", + filename=err_filename, + lineno=err_lineno, + cause=name_error, + suggestion=( + "Variables used inside staged control flow " + "(for/if/while) must be defined before the " + "control flow region. Give the variable an " + "initial value before the loop or branch." + ), + ) from name_error raise DSLUserCodeError( - f"NameError in `{funcBody.__name__}`: {name_error}", + DiagId.SCOPE_UNBOUND_NAME_IN_TRACE, filename=err_filename, lineno=err_lineno, cause=name_error, - suggestion=( - "Variables used inside staged control flow " - "(for/if/while) must be defined before the " - "control flow region. Give the variable an " - "initial value before the loop or branch." - ), + var=getattr(name_error, "name", None) + or f"", ) from name_error except (DSLRuntimeError, DSLUserCodeError): raise @@ -2236,6 +2333,8 @@ def generate_mlir( return result # Get a single reference to the cache since garbage collection + # This module-hash lookup needs no specialization validation: + # tracing re-ran before the hash was computed (self-curing). cached_jit_func = None if no_cache else self.jit_cache.get(module_hash) if ( @@ -2259,6 +2358,12 @@ def generate_mlir( original_function_name=original_function_name, funcBody=funcBody, ) + # F-SPEC: the fresh handle adopts the specialization + # record its producing trace sealed (V-11 validates + # __call__ re-entries; a cached handle keeps its own). + from .pyir_runtime import _pyir_take_sealed_spec + + jit_function.seal_specialization(_pyir_take_sealed_spec()) else: # cache hit log().info( @@ -2320,6 +2425,26 @@ def _run_preprocessor_impl( self, original_function: Any, ) -> Any: + # Attestation gate (PyIR-mode fact): a generator/coroutine entry only + # mints its generator/coroutine object at call time -- the body would + # never trace -- so the PyIR rewrite refuses before attesting the + # function; the base rewrite compiles these entries exactly as before. + if self.envar.enable_pyir: + _entry_code = getattr(original_function, "__code__", None) + if inspect.isasyncgenfunction(original_function) or ( + inspect.iscoroutinefunction(original_function) + ): + raise DSLUserCodeError( + DiagId.UNSUP_ASYNC, + filename=getattr(_entry_code, "co_filename", None), + lineno=getattr(_entry_code, "co_firstlineno", None), + ) + if inspect.isgeneratorfunction(original_function): + raise DSLUserCodeError( + DiagId.UNSUP_YIELD, + filename=getattr(_entry_code, "co_filename", None), + lineno=getattr(_entry_code, "co_firstlineno", None), + ) function_name = original_function.__name__ self.funcBody = original_function log().info("Started preprocessing [%s]", function_name) @@ -2348,6 +2473,11 @@ def _run_preprocessor_impl( original_function._preprocessed_signature = ( self._preprocess_mode.current_signature() ) + # F-COVER: attest the rewrite on the function object itself (the + # value-borne carrier); the trace-entry check validates it. + _choke_ver = getattr(self.preprocessor, "choke_set_version", None) + if _choke_ver is not None: + original_function.__pyir_rewritten__ = _choke_ver return preprocessor_session.exec( original_function.__name__, @@ -2538,7 +2668,7 @@ def _prepare_compilation( pipeline = kwargs.pop("pipeline", None) gpu_module_attrs = kwargs.pop("gpu_module_attrs", {}) - no_cache = kwargs.pop("no_cache", False) + no_cache = kwargs.pop("no_cache", False) or self.envar.no_cache no_jit_engine = kwargs.pop("no_jit_engine", False) compile_only = kwargs.pop("compile_only", False) @@ -2644,6 +2774,38 @@ def _func_impl( with __import__( f"{__package__}.multi_stage_manager", fromlist=["_jit_scope"] )._jit_scope(): + # PyIR entry attestation (V-9) + trace-arg intake for the + # INLINE nested-jit call. + __import__( + f"{__package__}.pyir_runtime", + fromlist=["_pyir_assert_entry_attested"], + )._pyir_assert_entry_attested(funcBody, self) + _pyir_default_binds: Any = [] + try: + _pyir_rt = __import__( + f"{__package__}.pyir_runtime", + fromlist=["_pyir_register_trace_args"], + ) + _pyir_rt._pyir_register_trace_args(args, kwargs) + # Taken parameter defaults bind at the plain call below + # with no dispatcher in the way: intake their holders, + # stage their meta-leaf reads, and seal their F-SPEC rows. + _pyir_rt._pyir_register_taken_default_holders( + funcBody, args, kwargs + ) + _pyir_default_binds = ( + _pyir_rt._pyir_boundary_stage_taken_default_reads( + funcBody, args, kwargs + ) + ) + _pyir_rt._pyir_spec_boundary_default_reads(funcBody, args, kwargs) + except Exception: + pass + if _pyir_default_binds: + try: + return funcBody(*args, **kwargs) + finally: + _pyir_rt._pyir_boundary_restore_meta_reads(_pyir_default_binds) return funcBody(*args, **kwargs) setup = self._prepare_compilation(funcBody, *args, **kwargs) @@ -2949,6 +3111,21 @@ def kernel_wrapper(*args: Any, **kwargs: Any) -> Any: "kernelGenHelper should be explicitly specified!" ) + # Taken parameter defaults bind at apply_defaults below with + # no dispatcher in the way: intake their holders and seal + # their F-SPEC rows. + try: + _pyir_rt = __import__( + f"{__package__}.pyir_runtime", + fromlist=["_pyir_spec_boundary_default_reads"], + ) + _pyir_rt._pyir_register_taken_default_holders( + funcBody, args, kwargs + ) + _pyir_rt._pyir_spec_boundary_default_reads(funcBody, args, kwargs) + except Exception: + pass + # Get bound arguments bound_args = self._get_function_bound_args( signature, kernel_name, *args, **kwargs @@ -3013,6 +3190,20 @@ def kernel_wrapper(*args: Any, **kwargs: Any) -> Any: f"{__package__}.multi_stage_manager", fromlist=["isolated_region"], ).isolated_region(): + # PyIR entry attestation (V-9), then trace-arg + # intake: kernel block-arg reconstructions become + # holders. + __import__( + f"{__package__}.pyir_runtime", + fromlist=["_pyir_assert_entry_attested"], + )._pyir_assert_entry_attested(funcBody, self) + try: + __import__( + f"{__package__}.pyir_runtime", + fromlist=["_pyir_register_trace_args"], + )._pyir_register_trace_args(ir_args, ir_kwargs) + except Exception: + pass kernel_ret = funcBody(*ir_args, **ir_kwargs) if hasattr(helper, "set_kernel_ret"): helper.set_kernel_ret(kernel_ret) diff --git a/python/CuTeDSL/cutlass/base_dsl/env_manager.py b/python/CuTeDSL/cutlass/base_dsl/env_manager.py index f2f7b3ae36..e78d8a5d0e 100644 --- a/python/CuTeDSL/cutlass/base_dsl/env_manager.py +++ b/python/CuTeDSL/cutlass/base_dsl/env_manager.py @@ -23,12 +23,15 @@ import sys import shutil import glob +import inspect import warnings +from dataclasses import dataclass, field from pathlib import Path -from functools import lru_cache +from functools import cache, lru_cache +from typing import Any, Callable, get_args from ..base_dsl.runtime.cuda import get_compute_capability_major_minor -from .common import DSLUserCodeError +from .common import DSLRuntimeError, DSLUserCodeError from .utils.logger import log from .cache_helpers import get_default_file_dump_root @@ -87,6 +90,192 @@ def _parse_keep_tokens(raw: str, prefix: str = "") -> frozenset[str]: return tokens - unknown +#: Superseded per-artifact switches, paired with the [DSL]_KEEP token each one +#: now folds into. Kept so existing scripts keep working for a release. +_DEPRECATED_KEEP_SWITCHES: tuple[tuple[str, str], ...] = ( + ("KEEP_IR", "ir-debug"), + ("KEEP_PTX", "ptx"), + ("KEEP_CUBIN", "cubin"), + ("KEEP_SASS", "sass"), +) + + +def _default_dump_dir() -> str: + """Directory artifacts land in when ``[DSL]_DUMP_DIR`` is unset.""" + return str(get_default_file_dump_root()) + + +def _resolve_keep_tokens(prefix: str) -> frozenset[str]: + """Artifacts requested by ``[DSL]_KEEP``, with the deprecated switches folded in. + + Each superseded switch warns and contributes its token, so the rest of the + DSL only ever consults the token set. + """ + raw = get_str_env_var(f"{prefix}_KEEP", "") + tokens: set[str] = set(_parse_keep_tokens(raw, prefix) if raw else frozenset()) + for switch, token in _DEPRECATED_KEEP_SWITCHES: + if get_bool_env_var(f"{prefix}_{switch}", False): + warnings.warn( + f"{prefix}_{switch} is deprecated; use {prefix}_KEEP={token} instead.", + DeprecationWarning, + stacklevel=2, + ) + tokens.add(token) + return frozenset(tokens) + + + +@dataclass(frozen=True) +class EnvVar: + """One setting on an :class:`EnvironmentVarManager`. + + Written in the class body as ``attribute: type = env_var(...)``, so the + attribute, its type and where its value comes from are stated together and + exactly once. ``source`` is either the environment variable's suffix -- the + value is read from ``{prefix}_{source}`` -- or a function of the manager + that computes it, for a setting with no variable of its own. How to read + the environment follows from the annotation. + + ``affects_compile`` says whether the setting is part of the JIT cache key; + it is required, so a setting cannot be added without answering. + """ + + source: str | Callable[[Any], Any] + affects_compile: bool = field(kw_only=True) + default: Any = None + read_as: str | None = None + attribute: str = field(default="", init=False) + + def __post_init__(self) -> None: + if not isinstance(self.source, str) and ( + self.default is not None or self.read_as is not None + ): + raise DSLRuntimeError( + "A computed setting takes neither `default` nor `read_as`: it has no " + "environment variable to fall back from or to be read under." + ) + + def __set_name__(self, owner: type, name: str) -> None: + object.__setattr__(self, "attribute", name) + + @property + def key_name(self) -> str: + return self.read_as or self.attribute + + def resolve( + self, manager: Any, prefix: str, parser: Callable[..., Any] | None + ) -> Any: + """Value this setting takes for ``manager``.""" + if not isinstance(self.source, str): + return self.source(manager) + assert parser is not None + # A default that reads other settings is written as a function of the + # manager. There is no ambiguity to resolve: parser_for_type admits only + # bool, int and str, so a callable is never itself a legitimate default. + default = self.default(manager) if callable(self.default) else self.default + return parser(f"{prefix}_{self.source}", default) + + +def env_var( + source: str | Callable[[Any], Any], + *, + affects_compile: bool, + default: Any = None, + read_as: str | None = None, +) -> Any: + """Declare a setting, as the default of an annotated class attribute. + + ``source`` is the variable's suffix -- the value is read from + ``{prefix}_{source}`` -- or a function of the manager, for a setting + computed rather than read. Either way it may use anything declared above. + + Returns ``Any``, the way :func:`dataclasses.field` does, so the declaration + can carry the attribute's real type. ``default`` may itself be a function + of the manager when it depends on a setting declared above. + """ + return EnvVar( + source, + affects_compile=affects_compile, + default=default, + read_as=read_as, + ) + + +def _annotated_type(owner: type, attr: str) -> Any: + """Declared type of ``attr``, searched up ``owner``'s MRO.""" + for klass in owner.__mro__: + annotation = inspect.get_annotations(klass).get(attr) + if annotation is not None: + return annotation + raise DSLRuntimeError( + f"{owner.__name__}.{attr} has no type annotation, so the parser for its " + f"environment variable cannot be determined. Annotate it on the class." + ) + + +def parser_for_type(annotation: Any) -> Callable[..., Any]: + args = get_args(annotation) + optional = type(None) in args + base = next((a for a in args if a is not type(None)), annotation) + if base is bool: + return get_bool_env_var + if base is int: + return get_int_or_none_env_var if optional else get_int_env_var + if base is str: + return get_str_env_var + raise DSLRuntimeError( + f"No environment-variable reader for {annotation!r}. Give the setting a " + f"bool, int or str annotation, or compute it from the manager instead." + ) + + +class EnvVarSpec: + """Turns the settings declared in a class body into a spec. + + ``_ENV_VAR_SPEC`` is what a class declares plus everything it inherits, in + declaration order. A subclass redeclaring an attribute replaces the + inherited declaration and keeps its position. + """ + + _ENV_VAR_SPEC: tuple[EnvVar, ...] = () + + def __init_subclass__(cls, **kwargs: Any) -> None: + super().__init_subclass__(**kwargs) + composed = {entry.attribute: entry for entry in cls._ENV_VAR_SPEC} + composed.update( + {v.attribute: v for v in vars(cls).values() if isinstance(v, EnvVar)} + ) + cls._ENV_VAR_SPEC = tuple(composed.values()) + + def _apply_env_var_spec(self, prefix: str) -> None: + for entry in type(self)._ENV_VAR_SPEC: + parser = ( + parser_for_type(_annotated_type(type(self), entry.attribute)) + if isinstance(entry.source, str) + else None + ) + setattr(self, entry.attribute, entry.resolve(self, prefix, parser)) + + def cache_key_str(self) -> str: + """Return the settings that are part of the JIT cache key.""" + rendered = [] + for entry in sorted(self._ENV_VAR_SPEC, key=lambda e: e.key_name): + if not entry.affects_compile: + continue + value = getattr(self, entry.key_name) + if value is None: + continue + rendered.append(f"{entry.key_name}={_render_cache_key_value(value)};") + return "".join(rendered) + + +def _render_cache_key_value(value: object) -> str: + # Sorted, because frozenset iteration order is randomized per process. + if isinstance(value, (set, frozenset)): + return repr(tuple(sorted(value, key=repr))) + return repr(value) + + # ============================================================================= # Environment Variable Helpers # ============================================================================= @@ -342,31 +531,137 @@ def _find_cuda_home() -> str | None: -@lru_cache(maxsize=1) -def _find_nvdisasm_binary() -> str: - """Return absolute path to the bundled nvdisasm binary.""" +# Fallback minimum nvdisasm (CUDA Toolkit) version for SASS dumping, as +# (major, minor). Used only when the build-time CUDA version is unavailable +# (see _min_nvdisasm_version). +MIN_NVDISASM_VERSION: tuple[int, int] = (13, 3) + + +def _min_nvdisasm_version() -> tuple[int, int]: + """Minimum supported nvdisasm version for SASS dumping. + + The floor is the CUDA version the DSL was built with — the bundled + toolchain that produces the CUBIN. nvdisasm is versioned with the CUDA + Toolkit it ships in, so the two are directly comparable; an older + nvdisasm may not understand the CUBINs this toolchain emits, while a + newer one can always read them. There is deliberately no upper bound: + the cu12-built DSL, for example, ships with the 13.3 nvdisasm wheel + (the first version published on PyPI), which disassembles its CUBINs + fine. Falls back to the hardcoded pin floor when the build-time version + is unavailable (e.g. a DSL client that does not implement + _get_cuda_version). + """ + try: + from .version_info import CUDA_VERSION + except Exception: + return MIN_NVDISASM_VERSION + return (CUDA_VERSION.major, CUDA_VERSION.minor) + + + +def _nvdisasm_suggestion() -> str: + floor = _min_nvdisasm_version() + return "\n".join( + [ + f"SASS dumping requires nvdisasm >= {floor[0]}.{floor[1]}." + " Any of the following works:", + " • pip install nvidia-cutlass-dsl[sass]", + " • install or upgrade a local CUDA Toolkit and expose it via" + " CUDA_HOME/CUDA_PATH", + ] + ) + + + +def _nvdisasm_from_wheel() -> str | None: from importlib import metadata try: dist = metadata.distribution("nvidia-cuda-nvdisasm") except metadata.PackageNotFoundError: - dist = None - if dist is not None and dist.files is not None: - for entry in dist.files: - if entry.name == "nvdisasm": - binpath = Path(str(dist.locate_file(entry))) - if binpath.is_file(): - return str(binpath) - raise DSLUserCodeError( - "nvdisasm binary not found inside the nvidia-cuda-nvdisasm wheel.", - suggestion="\n".join( - [ - "nvidia-cuda-nvdisasm is a runtime dependency of nvidia-cutlass-dsl", - "and should have been installed automatically. Try:", - " • pip install --force-reinstall nvidia-cuda-nvdisasm", - ] - ), - ) + return None + if dist.files is None: + return None + for entry in dist.files: + if entry.name == "nvdisasm": + binpath = Path(str(dist.locate_file(entry))) + if binpath.is_file(): + return str(binpath) + return None + + +def _nvdisasm_from_cuda_toolkit() -> str | None: + """Probe the CUDA Toolkit discovered by ``_find_cuda_home`` (CUDA_HOME / + CUDA_PATH, the root derived from nvcc on PATH, or a common install + location such as /usr/local/cuda*). + + Returns None when no toolkit is found or the toolkit has no nvdisasm. + """ + root = _find_cuda_home() + if not root: + return None + name = "nvdisasm.exe" if IS_WINDOWS else "nvdisasm" + binpath = Path(root) / "bin" / name + return str(binpath) if binpath.is_file() else None + + +def _get_nvdisasm_version(binary: str) -> tuple[int, int] | None: + """Return the (major, minor) CUDA version reported by ``nvdisasm --version``.""" + import re + import subprocess + + try: + result = subprocess.run( + [binary, "--version"], capture_output=True, text=True, check=True + ) + except (OSError, subprocess.CalledProcessError): + return None + # e.g. "Cuda compilation tools, release 13.5, V13.5.0" + match = re.search(r"release (\d+)\.(\d+)", result.stdout) + if match is None: + return None + return (int(match.group(1)), int(match.group(2))) + + +@cache +def _find_nvdisasm_binary() -> str: + """Locate a compatible nvdisasm binary for SASS dumping. + + Probe order (first hit wins): + 1. the nvidia-cuda-nvdisasm pip wheel (installed via the [sass] extra) + 2. the CUDA Toolkit located by _find_cuda_home (CUDA_HOME / CUDA_PATH, + the root derived from nvcc on PATH, or /usr/local/cuda*) + + The wheel is probed before the local toolkit so that users who installed + the [sass] extra get a predictable version regardless of local CTK state. + + The minimum supported version is derived from the CUDA version the DSL + was built with (see _min_nvdisasm_version); an incompatible version is a + hard error. + """ + binary = None + from_env_var = False + if binary is None: + binary = _nvdisasm_from_wheel() or _nvdisasm_from_cuda_toolkit() + if binary is None: + raise DSLUserCodeError( + "SASS dumping requires the nvdisasm tool, but it was not found.", + suggestion=_nvdisasm_suggestion(), + ) + version = _get_nvdisasm_version(binary) + floor = _min_nvdisasm_version() + if version is None or version < floor: + found = ( + "an unknown version" + if version is None + else f"version {version[0]}.{version[1]}" + ) + raise DSLUserCodeError( + f"nvdisasm at {binary!r} reports {found}, which is not supported" + " for SASS dumping.", + suggestion=_nvdisasm_suggestion(), + ) + return binary def dump_sass( @@ -409,6 +704,7 @@ def _get_libs_cand(start: str | Path) -> str | None: try: from .version_info import CUDA_VERSION + major = CUDA_VERSION.major lib_folder_guesses.append(f"cu{major}/lib") except Exception: @@ -463,16 +759,21 @@ def get_prefix_dsl_libs(prefix: str) -> str | None: return None -class LogEnvironmentManager: +class LogEnvironmentManager(EnvVarSpec): + jit_time_profiling: bool = env_var( + "JIT_TIME_PROFILING", affects_compile=False, default=False + ) + log_to_console: bool = env_var( + "LOG_TO_CONSOLE", affects_compile=False, default=False + ) + log_to_file: bool = env_var("LOG_TO_FILE", affects_compile=False, default=False) + log_level: int = env_var("LOG_LEVEL", affects_compile=False, default=1) + def __init__(self, prefix: str = "DSL") -> None: self.prefix = prefix - # Logging options - self.jit_time_profiling = get_bool_env_var( - f"{prefix}_JIT_TIME_PROFILING", False - ) - self.log_to_console = get_bool_env_var(f"{prefix}_LOG_TO_CONSOLE", False) - self.log_to_file = get_bool_env_var(f"{prefix}_LOG_TO_FILE", False) + self._apply_env_var_spec(prefix) + if ( has_env_var(f"{prefix}_LOG_LEVEL") and not self.log_to_console @@ -484,7 +785,6 @@ def __init__(self, prefix: str = "DSL") -> None: prefix, prefix, ) - self.log_level = get_int_env_var(f"{prefix}_LOG_LEVEL", 1) class EnvironmentVarManager(LogEnvironmentManager): @@ -555,122 +855,137 @@ class EnvironmentVarManager(LogEnvironmentManager): """ + # Master switch for DSL developers: raises the default of a curated set of + # diagnostic settings below, each still overridable by its own variable. One of + # those is lineinfo, which reaches GenerateLineInfo; it also injects + # warnings{nvvm} and alters MLIR locations. + debug: bool = env_var("DEBUG", affects_compile=True, default=False) + print_after_preprocessor: bool = env_var( + "PRINT_AFTER_PREPROCESSOR", affects_compile=False, default=False + ) + print_ir: bool = env_var("PRINT_IR", affects_compile=False, default=False) + # Selects between a full traceback and the formatted message on the exception + # path. Supplies the default of filter_stacktrace, never of lineinfo -- that + # coupling runs through debug. + show_stacktrace: bool = env_var( + "SHOW_STACKTRACE", affects_compile=False, default=lambda mgr: mgr.debug + ) + # Decides whether the frame-filtering excepthook is installed. Defaulted off + # under debug or show_stacktrace so internal DSL frames stay visible. + filter_stacktrace: bool = env_var( + "FILTER_STACKTRACE", + affects_compile=False, + default=lambda mgr: not (mgr.debug or mgr.show_stacktrace), + ) + enable_pyir: bool = env_var("ENABLE_PYIR", affects_compile=True, default=False) + auto_m2s: bool = env_var("AUTO_M2S", affects_compile=True, default=False) + tolerate_m2m: bool = env_var("TOLERATE_M2M", affects_compile=True, default=True) + lineinfo: bool = env_var( + "LINEINFO", affects_compile=True, default=lambda mgr: mgr.debug + ) + # Governs whether results are cached, not what is compiled. + no_cache: bool = env_var("NO_CACHE", affects_compile=False, default=False) + jit_cache_max_elems: int | None = env_var( + "JIT_CACHE_MAX_ELEMS", affects_compile=False, default=None + ) + # Chooses where artifacts are written. Reaches the compiler only through the + # keep_* dump paths, and those force caching off. + dump_dir: str = env_var( + "DUMP_DIR", + affects_compile=False, + default=lambda mgr: _default_dump_dir(), + ) + # Unread; the cache root is resolved separately from the environment by + # get_default_generated_ir_path. + cache_dir: str | None = env_var("CACHE_DIR", affects_compile=False, default=None) + # Every artifact the tokens below request is dumped either from build_module, + # which runs before the cache is consulted, or from a path that forces + # no_cache -- so none of them can be served stale, and the raw token set does + # not belong in the key either. A token wired up later answers for itself: + # the flag deriving it has to state its own affects_compile. + keep_tokens: frozenset[str] = env_var( + lambda mgr: _resolve_keep_tokens(mgr.prefix), affects_compile=False + ) + # Saves IR after canonicalize+cse, the readable form. + keep_ir_clean: bool = env_var( + lambda mgr: "ir" in mgr.keep_tokens, affects_compile=False + ) + # Saves raw IR before any passes, the old KEEP_IR=1 semantics. + keep_ir: bool = env_var( + lambda mgr: "ir-debug" in mgr.keep_tokens, affects_compile=False + ) + keep_ptx: bool = env_var( + lambda mgr: "ptx" in mgr.keep_tokens, affects_compile=False + ) + keep_cubin: bool = env_var( + lambda mgr: "cubin" in mgr.keep_tokens, affects_compile=False + ) + keep_sass: bool = env_var( + lambda mgr: "sass" in mgr.keep_tokens, affects_compile=False + ) + dryrun: bool = env_var("DRYRUN", affects_compile=True, default=False) + # Stored under _arch because the public spelling belongs to the arch property, + # which detects lazily; read_as sends the key through that property so a + # not-yet-detected architecture cannot drop out of it. + _arch: str | None = env_var( + "ARCH", affects_compile=True, default=None, read_as="arch" + ) + # Installs a warnings.filterwarnings("error") that turns a warning raised + # during compilation into an exception. Keyed because an artifact compiled + # while warnings were tolerated may be one that this setting is meant to + # reject, and a cache hit skips the compile that would have rejected it. + warnings_as_errors: bool = env_var( + "WARNINGS_AS_ERRORS", affects_compile=True, default=False + ) + warnings_ignore: bool = env_var( + "WARNINGS_IGNORE", affects_compile=False, default=False + ) + enable_optimization_warnings: bool = env_var( + "ENABLE_OPTIMIZATION_WARNINGS", + affects_compile=False, + default=lambda mgr: mgr.debug, + ) + # Governs whether results are cached, not what is compiled. + disable_file_caching: bool = env_var( + "DISABLE_FILE_CACHING", affects_compile=False, default=False + ) + compiler_opt: str = env_var("COMPILER_OPT", affects_compile=True, default="") + compiler_backend: str = env_var( + "COMPILER_BACKEND", affects_compile=True, default="legacy" + ) + # MLIR runtime libraries linked by the JIT. + shared_libs: str | None = env_var( + lambda mgr: get_prefix_dsl_libs(mgr.prefix), affects_compile=True + ) + # Enables asserts in host and device code. + enable_assertions: bool = env_var( + "ENABLE_ASSERTIONS", affects_compile=True, default=False + ) + enable_tvm_ffi: bool = env_var( + "ENABLE_TVM_FFI", affects_compile=True, default=False + ) + loc_tracebacks: int = env_var("LOC_TRACEBACKS", affects_compile=True, default=0) + def __init__(self, prefix: str = "DSL") -> None: super().__init__(prefix) - # Master debug switch for DSL developers. When True, it raises the - # default of a curated set of diagnostic/correctness settings below - # (lineinfo, stacktrace, optimization warnings, IR verification). - # Each of those settings remains independently overridable by its own - # env var, so debugging mode only changes their defaults. - self.debug = get_bool_env_var(f"{prefix}_DEBUG", False) + # PyIR-mode fact per DSL prefix: the verify-failure funnel names a + # mixed-flag configuration when a cross-DSL compile trips dominance. + try: + from .pyir_state import _pyir_register_mode_fact + + _pyir_register_mode_fact(prefix, self.enable_pyir) + except ImportError: + pass - # Printing options - self.print_after_preprocessor = get_bool_env_var( - f"{prefix}_PRINT_AFTER_PREPROCESSOR", False - ) - self.print_ir = get_bool_env_var(f"{prefix}_PRINT_IR", False) - # SHOW_STACKTRACE (and DEBUG) show the full, unfiltered traceback, so - # internal-frame filtering is disabled by default in either mode. - self.show_stacktrace = get_bool_env_var(f"{prefix}_SHOW_STACKTRACE", self.debug) - self.filter_stacktrace = get_bool_env_var( - f"{prefix}_FILTER_STACKTRACE", not (self.debug or self.show_stacktrace) - ) - self.enable_pyir = get_bool_env_var(f"{prefix}_ENABLE_PYIR", False) - self.auto_m2s = get_bool_env_var(f"{prefix}_AUTO_M2S", False) - self.tolerate_m2m = get_bool_env_var(f"{prefix}_TOLERATE_M2M", True) - - self.lineinfo = get_bool_env_var(f"{prefix}_LINEINFO", self.debug) - self.no_cache = get_bool_env_var(f"{prefix}_NO_CACHE", False) - self.jit_cache_max_elems = get_int_or_none_env_var( - f"{prefix}_JIT_CACHE_MAX_ELEMS", None - ) if self.no_cache: self.jit_cache_max_elems = 0 - self.dump_dir = get_str_env_var( - f"{prefix}_DUMP_DIR", str(get_default_file_dump_root()) - ) - # File options - self.cache_dir = get_str_env_var(f"{prefix}_CACHE_DIR", None) - - # ------------------------------------------------------------------ # - # Artifact keep — [DSL]_KEEP= # - # ------------------------------------------------------------------ # - # Parse new consolidated option. - _keep_raw = get_str_env_var(f"{prefix}_KEEP", "") - _keep_tokens: set[str] = set( - _parse_keep_tokens(_keep_raw, prefix) if _keep_raw else frozenset() - ) - # Backward compatibility: publicly-documented old options emit a - # DeprecationWarning and fold into _keep_tokens. - if get_bool_env_var(f"{prefix}_KEEP_IR", False): - warnings.warn( - f"{prefix}_KEEP_IR is deprecated; use {prefix}_KEEP=ir-debug instead.", - DeprecationWarning, - stacklevel=2, - ) - _keep_tokens.add("ir-debug") - if get_bool_env_var(f"{prefix}_KEEP_PTX", False): - warnings.warn( - f"{prefix}_KEEP_PTX is deprecated; use {prefix}_KEEP=ptx instead.", - DeprecationWarning, - stacklevel=2, - ) - _keep_tokens.add("ptx") - if get_bool_env_var(f"{prefix}_KEEP_CUBIN", False): - warnings.warn( - f"{prefix}_KEEP_CUBIN is deprecated; use {prefix}_KEEP=cubin instead.", - DeprecationWarning, - stacklevel=2, - ) - _keep_tokens.add("cubin") - - if get_bool_env_var(f"{prefix}_KEEP_SASS", False): - warnings.warn( - f"{prefix}_KEEP_SASS is deprecated; use {prefix}_KEEP=sass instead.", - DeprecationWarning, - stacklevel=2, - ) - _keep_tokens.add("sass") - self.keep_tokens: frozenset[str] = frozenset(_keep_tokens) - - # Derived boolean attributes — used by compiler.py and dsl.py. - # keep_ir_clean: save IR after canonicalize+cse (the readable form). - self.keep_ir_clean: bool = "ir" in self.keep_tokens - # keep_ir: save raw IR before any passes (old KEEP_IR=1 semantics). - self.keep_ir: bool = "ir-debug" in self.keep_tokens - self.keep_ptx: bool = "ptx" in self.keep_tokens - self.keep_cubin: bool = "cubin" in self.keep_tokens - self.keep_sass: bool = "sass" in self.keep_tokens - _check_nvdisasm_wheel = True - if _check_nvdisasm_wheel and self.keep_sass: + # Fail at construction (rather than after a long compile) when the + # user asked for SASS dumping but no usable nvdisasm is available. + _check_nvdisasm = True + if _check_nvdisasm and self.keep_sass: _find_nvdisasm_binary() - # Other options - self.dryrun = get_bool_env_var(f"{prefix}_DRYRUN", False) - self._arch: str | None = get_str_env_var(f"{prefix}_ARCH") - self.warnings_as_errors = get_bool_env_var( - f"{prefix}_WARNINGS_AS_ERRORS", False - ) - self.warnings_ignore = get_bool_env_var(f"{prefix}_WARNINGS_IGNORE", False) - self.enable_optimization_warnings = get_bool_env_var( - f"{prefix}_ENABLE_OPTIMIZATION_WARNINGS", self.debug - ) - self.disable_file_caching = get_bool_env_var( - f"{prefix}_DISABLE_FILE_CACHING", False - ) - self.compiler_opt = get_str_env_var(f"{prefix}_COMPILER_OPT", "") - self.compiler_backend = get_str_env_var(f"{prefix}_COMPILER_BACKEND", "legacy") - - # set mlir shared libraries - self.shared_libs = get_prefix_dsl_libs(prefix) - - # whether to enable assert in host and device code - self.enable_assertions = get_bool_env_var(f"{prefix}_ENABLE_ASSERTIONS", False) - - self.enable_tvm_ffi = get_bool_env_var(f"{prefix}_ENABLE_TVM_FFI", False) - - self.loc_tracebacks = get_int_env_var(f"{prefix}_LOC_TRACEBACKS", 0) @property def arch(self) -> str: diff --git a/python/CuTeDSL/cutlass/base_dsl/jit_executor.py b/python/CuTeDSL/cutlass/base_dsl/jit_executor.py index 62e4d2ac03..2ec405f211 100644 --- a/python/CuTeDSL/cutlass/base_dsl/jit_executor.py +++ b/python/CuTeDSL/cutlass/base_dsl/jit_executor.py @@ -15,7 +15,7 @@ Pointer-address runtime arguments are opt-in. A compile-time example such as ``cute.runtime.nullptr(dtype, space)`` marks the corresponding runtime argument as a pointer slot, and -``ExecutionArgs.record_pointer_arg_specs_from_compile_args`` records that +``ExecutionArgs.record_arg_specs_from_compile_args`` records that metadata in ``_pointer_address_arg_specs``. Runtime calls may then pass an integer address, ``ctypes.c_void_p``, a ctypes pointer object, or ``cute.runtime.nullptr(...)`` for a null address. Native JIT execution packs raw @@ -38,7 +38,18 @@ import ctypes import inspect import io -from typing import Any, NamedTuple, TYPE_CHECKING, ClassVar, cast, get_args, get_origin +from typing import ( + Any, + Literal, + NamedTuple, + TYPE_CHECKING, + ClassVar, + TypeVar, + cast, + get_args, + get_origin, + overload, +) from collections.abc import Callable, Sequence import weakref import threading @@ -378,7 +389,7 @@ class ExecutionArgs: Besides normal signature rectification and scalar casting, this class owns the pointer-address conversion contract. - ``record_pointer_arg_specs_from_compile_args`` derives a lightweight spec + ``record_arg_specs_from_compile_args`` derives a lightweight spec tree from compile-time arguments, and ``generate_execution_args`` / ``convert_python_pointer_args_for_tvm_ffi`` use that spec to decide where Python pointer-like values may replace runtime @@ -405,6 +416,9 @@ def __init__( None ] * self._meta.arg_count self._has_pointer_address_arg_specs = False + # Whether the function was compiled with ``None`` in each argument slot, + # or "unknown" when the compile-time arguments were not recorded. + self._compiled_with_none: list[Any] = ["unknown"] * self._meta.arg_count # When True (set when debugging mode is ON), generate_execution_args # runs thorough per-argument validation that is otherwise skipped to # keep the launch path fast. @@ -427,7 +441,7 @@ def has_pointer_address_arg_specs(self) -> bool: """Whether this compiled signature has pointer-address conversion slots.""" return self._has_pointer_address_arg_specs - def record_pointer_arg_specs_from_compile_args( + def record_arg_specs_from_compile_args( self, args: tuple[Any, ...], kwargs: dict[str, Any] ) -> None: """Record pointer conversion metadata from compile-time arguments. @@ -447,6 +461,12 @@ def record_pointer_arg_specs_from_compile_args( ``list[cutlass.Pointer]`` or ``Sequence[cutlass.Pointer]``, the inner annotation is applied to each element. + The same pass records ``_compiled_with_none``, marking which slots held + ``None``. Tracing takes ``None`` as a constexpr and folds it into the + kernel, so those slots carry no runtime argument; ``_validate_args_full`` + needs to tell them apart from slots the kernel really does expect an + argument for, which the annotation cannot say. + ``args``/``kwargs`` are first rectified through the function signature unless they already exactly match the positional runtime shape. That path may raise the usual argument-binding ``DSLRuntimeError`` cases: @@ -466,6 +486,7 @@ def record_pointer_arg_specs_from_compile_args( ) for index, arg in enumerate(input_args) ] + self._compiled_with_none = [arg is None for arg in input_args] self._has_pointer_address_arg_specs = any( spec is not None for spec in self._pointer_address_arg_specs ) @@ -633,6 +654,20 @@ def get_rectified_args( return rectified + def bound_call_arguments( + self, args: tuple[Any, ...], kwargs: dict[str, Any] + ) -> dict[str, Any]: + """THIS call's launch binding as name->value facts: the same tail-arg + slice and rectification ``_generate_execution_args`` performs, keyed + by the filtered runtime signature's parameter names.""" + n = self._meta.arg_count + head = args[:n] if len(args) > n else args + if not kwargs and len(head) == n: + input_args: Sequence[Any] = head + else: + input_args = self.get_rectified_args(head, kwargs) + return dict(zip(self._meta.all_names, input_args)) + def generate_execution_args( self, args: tuple[Any, ...], kwargs: dict[str, Any] ) -> tuple[list[Any], list[Any]]: @@ -787,6 +822,22 @@ def _validate_args_full(self, input_args: Sequence[Any]) -> None: name = meta.all_names[index] if index < len(meta.all_names) else f"#{index}" spec = self._pointer_address_arg_specs[index] + # A parameter compiled with `None` was baked in as a constexpr, so `None` is the only value + # it can take, and marshalling it to zero C pointers is correct. + compiled_with_none = self._compiled_with_none[index] + if compiled_with_none is True: + if arg is not None: + raise DSLUserCodeError( + DiagId.ARG_CONSTEXPR_MISMATCH, + arg_name=name, + arg_type=type(arg).__name__, + context={"position": index}, + ) + continue + + # Skip validation if not recorded. + if arg is None and compiled_with_none == "unknown": + continue # Numeric scalars must be number-like so the later t.cast succeeds; # only reject clearly non-numeric values to avoid false positives on # the DSL's own numeric wrapper types. @@ -1071,6 +1122,21 @@ def __del__(self) -> None: self.unload() +def _pyir_validate_spec_reentry( + spec: "tuple | None", + execution_args: "ExecutionArgs | None", + args: tuple, + kwargs: dict, +) -> None: + """Lazy hook into the F-SPEC re-entry engine (pyir_spec); a handle with + no sealed spec never imports the engine.""" + if spec is None: + return + from .pyir_spec import _pyir_validate_spec_reentry as _validate + + _validate(spec, execution_args, args, kwargs) + + class JitExecutor: """An executable function that can be called to launch a device kernel. @@ -1083,12 +1149,16 @@ def __init__( jit_module: JitModule | Any, exec_context: JitExecuteContext | None, jit_time_profiling: bool, + spec: "tuple | None" = None, ) -> None: # JitExecutor will keep JitCompiledFunction alive so that the underlying # ExecutionEngine and module data is not discarded until runtime callables # are garbage collected. self.jit_module = jit_module self.exec_context = exec_context + # F-SPEC: (record, complete, entry_func, receiver) sealed by the trace that + # produced this module; validated at every __call__ re-entry. + self._pyir_spec = spec self.profiler = timer(enable=jit_time_profiling) if jit_time_profiling else None # Get the cuda result type from the capi function. @@ -1172,13 +1242,22 @@ def run_compiled_program(self, exe_args: list[Any]) -> int | None: error_code = self.cuda_result.value # type: ignore[union-attr] if error_code == 0: return error_code - raise cuda_helpers.create_cuda_runtime_error(error_code) + raise cuda_helpers.create_cuda_runtime_error( + error_code, cuda_helpers.cudart.cudaError_t + ) except DSLCudaRuntimeError as e: raise e except Exception as e: raise DSLRuntimeError(f"💥💥💥 Runtime Crash 💥💥💥", cause=e) def __call__(self, *args: Any, **kwargs: Any) -> int | None: + if self._pyir_spec is not None: + _pyir_validate_spec_reentry( + self._pyir_spec, + getattr(self.jit_module, "execution_args", None), + args, + kwargs, + ) exe_args, adapted_args = self.generate_execution_args(*args, **kwargs) return self.run_compiled_program(exe_args) @@ -1251,6 +1330,9 @@ def __init__(self, func_ptr: int, execution_args: ExecutionArgs) -> None: """ +AuxRuntimeFuncT = TypeVar("AuxRuntimeFuncT", bound=AuxRuntimeFunc) + + class JitCompiledFunction: """Holds a compiled function.""" @@ -1293,7 +1375,7 @@ def __init__( full_arg_check=full_arg_check, adapter_scope=self._jit_arg_adapter_scope, ) - self.execution_args.record_pointer_arg_specs_from_compile_args( + self.execution_args.record_arg_specs_from_compile_args( tuple(dynamic_args or ()), dynamic_kwargs or {} ) self.jit_time_profiling = jit_time_profiling @@ -1323,6 +1405,11 @@ def __init__( self._executor_lock = threading.RLock() self._default_executor: JitExecutor | None = None + # F-SPEC: (record, complete, entry_func, receiver) sealed by the producing + # trace; validated at __call__ re-entry on this handle and propagated + # to every executor derived from it. + self._pyir_spec: "tuple | None" = None + # This is used to do early generation of the c header arguments to release the reference to the dynamic arguments. self._generate_c_header_arguments(dynamic_args, dynamic_kwargs) @@ -1410,34 +1497,65 @@ def to(self, device: Any = None) -> JitExecutor: # Create a new executor that will be tied to a device context # n.b. host only modules do not load device specific modules or context. context = self.jit_module.get_device_execute_context(device) - return JitExecutor(self.jit_module, context, self.jit_time_profiling) + return JitExecutor( + self.jit_module, + context, + self.jit_time_profiling, + spec=self._pyir_spec, + ) def generate_execution_args( self, *args: Any, **kwargs: Any ) -> tuple[list[Any], list[Any]]: return self.execution_args.generate_execution_args(args, kwargs) + @overload + def get_aux_func( + self, + func_class: type[AuxRuntimeFuncT], + kernel: Callable[..., Any] | None = None, + *, + required: Literal[True] = True, + ) -> AuxRuntimeFuncT: ... + + @overload + def get_aux_func( + self, + func_class: type[AuxRuntimeFuncT], + kernel: Callable[..., Any] | None = None, + *, + required: Literal[False], + ) -> AuxRuntimeFuncT | None: ... + def get_aux_func( - self, func_class: type[AuxRuntimeFunc], kernel: Callable[..., Any] - ) -> AuxRuntimeFunc: - """Look up and return an auxiliary runtime function for a specific kernel. + self, + func_class: type[AuxRuntimeFuncT], + kernel: Callable[..., Any] | None = None, + *, + required: bool = True, + ) -> AuxRuntimeFuncT | None: + """Look up and return an auxiliary runtime function. - ``kernel`` must be a ``@dsl_name.kernel``-annotated callable that was called - inside the ``@dsl_name.jit`` function that produced this compiled object. - The lookup resolves the symbol ``{kernel_name}_{func_class.name}`` for - that specific kernel. + When ``kernel`` is provided, resolves + ``{kernel_name}_{func_class.name}``. Otherwise, resolves + ``func_class.name`` directly. :param func_class: A subclass of :class:`AuxRuntimeFunc` whose - ``name`` class attribute identifies the host function suffix. - :param kernel: A ``@dsl_name.kernel``-annotated callable. Must have been - called at least once inside a ``@dsl_name.jit`` function so that - ``_dsl_kernel_name`` is set. + ``name`` class attribute identifies the full symbol when ``kernel`` + is omitted or the host function suffix when ``kernel`` is provided. + :param kernel: An optional ``@dsl_name.kernel``-annotated callable. If + provided, it must have been called at least once inside a + ``@dsl_name.jit`` function so that ``_dsl_kernel_name`` is set. + :param required: Whether an unavailable auxiliary function is an error. + If false, a missing execution engine or symbol returns ``None``. :return: An instance of ``func_class`` initialised with the matched - function pointer and ready to call. + function pointer, or ``None`` when the symbol is not required and + is absent. :raises TypeError: If ``func_class`` is not a subclass of :class:`AuxRuntimeFunc`. :raises ValueError: If ``kernel`` has no ``_dsl_kernel_name`` attribute. - :raises DSLRuntimeError: If no matching symbol is found in the JIT engine. + :raises DSLRuntimeError: If ``required`` is true and the execution + engine is unavailable or no matching symbol is found in it. """ if not ( isinstance(func_class, type) and issubclass(func_class, AuxRuntimeFunc) @@ -1446,20 +1564,26 @@ def get_aux_func( f"func_class must be a subclass of AuxRuntimeFunc, got {func_class!r}" ) - # Unwrap bound methods then @wraps wrappers to reach the original funcBody. - func_body = getattr(kernel, "__func__", kernel) # bound method → function - func_body = getattr(func_body, "__wrapped__", func_body) # jit_wrapper → func - kernel_name = getattr(func_body, "_dsl_kernel_name", None) - if kernel_name is None: - raise ValueError( - f"kernel {kernel!r} has no '_dsl_kernel_name' attribute. " - "Make sure it has been called at least once inside a @cute.jit function." - ) + sym_name = func_class.name + candidate = sym_name + if kernel is not None: + # Unwrap bound methods then @wraps wrappers to reach the original funcBody. + func_body = getattr(kernel, "__func__", kernel) + func_body = getattr(func_body, "__wrapped__", func_body) + kernel_name = getattr(func_body, "_dsl_kernel_name", None) + if kernel_name is None: + raise ValueError( + f"kernel {kernel!r} has no '_dsl_kernel_name' attribute. " + "Make sure it has been called at least once inside a " + "@cute.jit function." + ) + candidate = f"{kernel_name}_{sym_name}" + + if self.engine is None and not required: + return None self._validate_engine() - sym_name = func_class.name - candidate = f"{kernel_name}_{sym_name}" candidates = [candidate] if self.prefix is not None: candidates = [f"_mlir_{self.prefix}_{candidate}"] + candidates @@ -1471,12 +1595,19 @@ def get_aux_func( break if not fn_ptr: + if not required: + return None raise DSLRuntimeError( f"Host function '{sym_name}' not found in JIT engine. " f"Tried: {candidates}" ) return func_class(fn_ptr, self.execution_args) + def seal_specialization(self, spec: "tuple | None") -> None: + """Attach the producing trace's F-SPEC (record, complete, entry_func, receiver); + __call__ re-entries on this handle validate against it (V-11).""" + self._pyir_spec = spec + def __call__(self, *args: Any, **kwargs: Any) -> int | None: """Executes the jit-compiled function under the currently active CUDA context. @@ -1484,6 +1615,13 @@ def __call__(self, *args: Any, **kwargs: Any) -> int | None: CUDA errors. If you need to call the kernel on multiple devices use `to` to return a per-device function. """ + if self._pyir_spec is not None: + _pyir_validate_spec_reentry( + self._pyir_spec, + getattr(self, "execution_args", None), + args, + kwargs, + ) exe_args, adapted_args = self.execution_args.generate_execution_args( args, kwargs ) diff --git a/python/CuTeDSL/cutlass/base_dsl/multi_stage_manager.py b/python/CuTeDSL/cutlass/base_dsl/multi_stage_manager.py index 4765cf87b5..4610a99108 100644 --- a/python/CuTeDSL/cutlass/base_dsl/multi_stage_manager.py +++ b/python/CuTeDSL/cutlass/base_dsl/multi_stage_manager.py @@ -16,6 +16,8 @@ """ +import itertools + from collections.abc import Iterator from contextlib import contextmanager from typing import Any @@ -50,24 +52,32 @@ def is_inside_staged_cf() -> bool: return _staged_cf_depth > 0 -# ``range_constexpr`` / ``const_expr`` if / ``const_expr`` while resolve at -# trace time (the loop unrolls or the branch is selected in Python), so a -# Python-side ``.append`` / variable mutation directly inside their body is -# realized deterministically -- it is NOT a silently-lost loop-carried -# mutation. The M->M / container-mutation guards therefore opt out when the -# INNERMOST control-flow construct governing the mutation is a constexpr -# construct. -# -# "Innermost" is the key: a constexpr scope nested inside dynamic staged CF -# still opts out (the mutation is constexpr-governed), but a *dynamic* loop / -# if nested inside a constexpr scope must NOT opt out (the mutation is -# governed by the dynamic construct and is silently lost). We capture this -# by recording ``_staged_cf_depth`` at each constexpr-scope entry: the scope -# is the innermost CF only while no dynamic CF has been entered since, i.e. -# the current ``_staged_cf_depth`` still equals the snapshot on top of the -# stack. +def current_staged_cf_depth() -> int: + """Current staged-CF nesting depth (0 outside any region): a first-def at a + shallower depth than a reassignment marks a loop-carried/cross-branch variable.""" + return _staged_cf_depth + + +# Each entry records ``_staged_cf_depth`` at constexpr-scope entry; the scope is +# the innermost CF only while the current depth still equals the top snapshot. _constexpr_scope_stack: list[int] = [] +# Monotone per-INSTANCE serials, parallel to the depth stack: each +# ``enter_constexpr_loop()`` (one unrolled iteration / one taken const_expr +# arm) is a distinct instance; binding-birth ownership records these serials. +_constexpr_instance_counter = itertools.count(1) +_constexpr_instance_serials: list[int] = [] + + +def open_constexpr_instance_serial() -> "int | None": + """Serial of the innermost OPEN constexpr instance, else ``None``.""" + return _constexpr_instance_serials[-1] if _constexpr_instance_serials else None + + +def open_constexpr_instance_serials() -> "list[int]": + """All open constexpr-instance serials (outermost first).""" + return _constexpr_instance_serials + def enter_constexpr_loop() -> None: """Enter a constexpr scope (``range_constexpr`` body or a @@ -79,12 +89,15 @@ def enter_constexpr_loop() -> None: the AST preprocessor instead emits a generated ``try/finally``. """ _constexpr_scope_stack.append(_staged_cf_depth) + _constexpr_instance_serials.append(next(_constexpr_instance_counter)) def exit_constexpr_loop() -> None: """Leave the innermost constexpr scope.""" if _constexpr_scope_stack: _constexpr_scope_stack.pop() + if _constexpr_instance_serials: + _constexpr_instance_serials.pop() def is_inside_constexpr_loop() -> bool: @@ -97,6 +110,12 @@ def is_inside_constexpr_loop() -> bool: ) +def constexpr_scope_under_staged_cf() -> bool: + """True if the innermost open constexpr scope was entered while staged CF was + open: it unrolls at trace time but its writes stay runtime-conditional.""" + return bool(_constexpr_scope_stack) and _constexpr_scope_stack[-1] > 0 + + def _reset_constexpr_scope() -> None: """Clear constexpr-scope state. Backstop for the AST-injected ``try/finally`` bracketing: cleared at outermost trace exit so an @@ -104,6 +123,7 @@ def _reset_constexpr_scope() -> None: silently disable the guards for the next trace. """ _constexpr_scope_stack.clear() + _constexpr_instance_serials.clear() @contextmanager @@ -119,22 +139,9 @@ def constexpr_loop_scope() -> Iterator[None]: exit_constexpr_loop() -def get_staged_cf_depth() -> int: - """Return the current staged-CF nesting depth. - - Counts runtime CF regions (``scf.if`` / ``scf.for`` / ``scf.while``) - entered via ``_scf_execute_pyir``. Trace-time ``range_constexpr`` - unrolls do NOT bump this counter, so a strictly greater depth between - a slot's first-def and a later reassignment means the reassignment is - inside a NESTED RUNTIME loop/if that must thread the value as an - iter_arg / scf result rather than constexpr-fold it. - """ - return _staged_cf_depth - - @contextmanager def _jit_scope() -> Iterator[None]: - """Track the staged-CF depth baseline for a nested ``@cute.jit`` call. + """Track the staged-CF depth baseline for a nested jit-decorated call. On exit from the OUTERMOST trace, clear the D1 meta-value promotion state so a second invocation of the same function in the same process @@ -149,16 +156,15 @@ def _jit_scope() -> Iterator[None]: # Outermost trace just exited -- clear D1 trace state and # any constexpr-scope state left over from an unbalanced trace. _reset_constexpr_scope() - try: - from .pyir_runtime import _exit_function_trace + # Function-local import (layering cycle: the pyir chain imports + # from this module at module top). + from .pyir_runtime import _exit_function_trace - _exit_function_trace() - except ImportError: - pass # PyIR runtime unavailable -- nothing to clean. + _exit_function_trace() def is_inside_locally_staged_cf() -> bool: - """Return True if the current ``@cute.jit`` body opened staged CF.""" + """Return True if the current jit-decorated body opened staged CF.""" if not _jit_depth_baseline_stack: return _staged_cf_depth > 0 return _staged_cf_depth > _jit_depth_baseline_stack[-1] @@ -193,23 +199,30 @@ def isolated_region() -> Iterator[None]: _staged_cf_depth = saved +# Lazily bound ``_WatchedM`` class (see :func:`_is_staged_value`). +_WATCHED_M_CLS: "type | None" = None + + def _is_staged_value(val: object) -> bool: """Return True if *val* is a staged (S) value (DSL type with ir_value). Bare Python scalars (int, float, bool, str, None) are Meta (M). DSL types (Numeric, Pointer, Array, etc.) are Staged (S). """ + global _WATCHED_M_CLS if val is None or isinstance(val, (int, float, bool, str, bytes, type)): return False - # Lazy import to avoid circular dependency: pyir_runtime imports from - # this module. - try: + # Lazily bound (layering cycle: the pyir chain imports from this module + # at module top); resolved once, then a plain global read. + wm = _WATCHED_M_CLS + if wm is None: from .pyir_runtime import _WatchedM - except ImportError: - _WatchedM = None # type: ignore[misc,assignment] - if _WatchedM is not None and isinstance(val, _WatchedM): + + wm = _WATCHED_M_CLS = _WatchedM + if isinstance(val, wm): return False - return hasattr(val, "ir_value") and callable(getattr(val, "ir_value")) + iv = getattr(val, "ir_value", None) + return iv is not None and callable(iv) def _get_ir_type(val: Any) -> Any: @@ -221,6 +234,84 @@ def _get_ir_type(val: Any) -> Any: return None +def _lift_scalar_to_staged_like(old_value: Any, new_value: Any) -> Any: + """Lift a Python scalar *new_value* to a staged value matching the staged + *old_value*'s type, or return ``None`` when no lifting applies. + + The lifting keys on the *capabilities* of the staged value's type, never + on its name, so it covers every staged scalar wrapper: + + 1. Numeric-typed wrappers expose an ``isinstance`` classmethod and a + scalar-accepting constructor -- lift via ``dsl_type(new_value)``. + + 2. Raw MLIR-value wrappers (``ir.Value`` subclasses without ``isinstance``) + lift to the ``Numeric`` of their MLIR element type. + """ + # Only a Python scalar can be lifted to a staged scalar; the guard lives + # here so the lifting rule is fully self-contained. + if not isinstance(new_value, (bool, int, float)): + return None + + # Function-local imports (layering cycle: typing/_mlir_helpers import back + # into base_dsl). + from .typing import Numeric + from .._mlir_helpers.arith import element_type + from .common import DSLRuntimeError + + dsl_type = type(old_value) + + # Strategy 1: Numeric-style type with an ``isinstance`` classmethod (bare + # ``ir.Value`` subclasses lack it, separating the two families by lookup). + type_isinstance = getattr(dsl_type, "isinstance", None) + if callable(type_isinstance): + try: + liftable = type_isinstance(new_value) + except (TypeError, ValueError): + liftable = False + if liftable: + try: + return dsl_type(new_value) + except (TypeError, AttributeError) as e: + # ``isinstance`` accepted the scalar but construction failed: + # an internal build failure, not a not-liftable slot. + raise DSLRuntimeError( + f"failed to lift scalar {new_value!r} into staged type " + f"{dsl_type.__name__}" + ) from e + # ``None`` = "no applicable lifting"; the caller raises the + # illegal-mutation diagnostic. + return None + + # Strategy 2: staged value backed by a raw MLIR value (no Numeric + # ``isinstance``). Lift to a Numeric of the value's element type. + ir_type = _get_ir_type(old_value) + if ir_type is not None: + try: + numeric_type = Numeric.from_mlir_type(element_type(ir_type)) + except (KeyError, ValueError, DSLRuntimeError): + # No Numeric for this element type: not liftable. + numeric_type = None + if numeric_type is not None: + try: + liftable = numeric_type.isinstance(new_value) + except (TypeError, ValueError): + liftable = False + if liftable: + try: + return numeric_type(new_value) + except (TypeError, AttributeError) as e: + # Construction failed after isinstance accepted: internal + # build failure, not a not-liftable slot. + raise DSLRuntimeError( + f"failed to lift scalar {new_value!r} into staged type " + f"{numeric_type.__name__}" + ) from e + + # ``None`` = "no applicable lifting"; the caller raises the + # illegal-mutation diagnostic. + return None + + def _user_type_name(value: Any) -> str: """A user-facing type name for *value*. @@ -247,8 +338,9 @@ def assign_meta_staged_check( target_name: str, old_value: Any, new_value: Any, - filename: str, - lineno: int, + filename: "str | None", + lineno: "int | None", + owner: Any = None, ) -> Any: """Runtime check for assignments inside staged control flow. @@ -274,13 +366,11 @@ def assign_meta_staged_check( # Rule 2: (M) mutation inside staged CF if old_is_staged and not new_is_staged: - # Auto-coerce: Python bool/int/float → matching DSL type - dsl_type = type(old_value) - if isinstance(new_value, (bool, int, float)) and dsl_type.isinstance(new_value): - try: - return dsl_type(new_value) - except (TypeError, ValueError): - pass + # Auto-lift a Python scalar reassigned to a staged variable back to + # staged, so the IR keeps an SSA value instead of dropping the write. + lifted_value = _lift_scalar_to_staged_like(old_value, new_value) + if lifted_value is not None: + return lifted_value raise DSLUserCodeError( DiagId.PHASE_ASSIGN_PYTHON_TO_TRACKED, @@ -290,23 +380,12 @@ def assign_meta_staged_check( ) if not old_is_staged and not new_is_staged: - if target_name == "self.dummy": - return None - - # D1 (META_VALUE_TABLE_DESIGN): when *old_value* is a ``_WatchedM`` - # wrapper, the slot is being tracked by the meta-value table. - # ``pyir_assign`` (below in pyir_runtime.py) handles retroactive - # promotion -- it creates a ``pyir.ref`` at function entry and - # rewrites baked constants via ``replaceAllUsesWith``. Allow the - # mutation here; D1 enforces correctness, including any type - # mismatch which surfaces as an MLIR verification error. - try: - from .pyir_runtime import _WatchedM + # A ``_WatchedM`` slot is table-tracked and ``pyir_assign`` handles its + # retroactive promotion — allow it here (local import: layering cycle). + from .pyir_runtime import _WatchedM - if isinstance(old_value, _WatchedM) or isinstance(new_value, _WatchedM): - return None - except ImportError: - pass + if isinstance(old_value, _WatchedM) or isinstance(new_value, _WatchedM): + return None # Mp→Mp: meta-primitive mutation inside staged CF. # When AUTO_M2S is enabled, auto-promote to staged so the @@ -325,8 +404,16 @@ def assign_meta_staged_check( except Exception: pass # fall through to error + # An adopted watched container is the same PLACE as the plain container + # that rebinds it -- adoption must not create a phase error. Same + # value-kind == same plain type under the transparency protocol + # (F-TRANSPARENT); no wrapper-class imports. + def _plain_type(v: Any) -> type: + return getattr(type(v), "__pyir_plain_type__", type(v)) + + _same_type = _plain_type(old_value) is _plain_type(new_value) same_type_compound = ( - type(old_value) is type(new_value) + _same_type and old_value is not None and not isinstance(old_value, (int, float, bool, str, bytes)) ) @@ -346,17 +433,51 @@ def assign_meta_staged_check( # per-field pyir_assign calls by pyir_assign(). # Only let them through if TOLERATE_M2M is True (default). if tolerate_m2m and same_type_compound: - try: - from .pyir_runtime import _has_decomposable_staged_fields + # Function-local import (layering cycle: the pyir chain imports + # from this module at module top). + from .pyir_runtime import _has_decomposable_staged_fields - if _has_decomposable_staged_fields(old_value): - return None - except ImportError: - pass # pyir_runtime not available (non-PyIR build) + if _has_decomposable_staged_fields(old_value): + return None + + # A mutation governed by a constexpr scope resolves at trace time and is + # meta by construction (the scope opens without raising staged-CF depth). + if is_inside_constexpr_loop(): + return None + + # A plain mutable ``set`` rebound over a set (``s = s | {x}``) bakes + # the trace-time snapshot: there is no watched-set wrapper and a set + # has no decomposable staged leaves, so the same-type tolerance below + # would admit it silently and every iteration would read one frozen + # union. Fail closed. ``frozenset`` rebinds stay admitted: their + # per-call-site bake is a pinned enables-lock + # (test_pyir_gap_frozenset_conditional_rebind_staged_if). + if isinstance(old_value, set) and isinstance(new_value, set): + raise DSLUserCodeError( + DiagId.PHASE_MUTATE_PYTHON, + filename=filename, + lineno=lineno, + var=target_name, + ) # Same-type M->M reassignment — tolerate if flag is set. if tolerate_m2m and same_type_compound: return None + + # In-place-mutation helper pattern (`x = obj.method(x)` returning None): + # state provably lives in the decomposable leaves, so dropping is benign. + if ( + new_value is None + and old_value is not None + and not isinstance(old_value, (int, float, bool, str, bytes)) + ): + # Function-local import (layering cycle: the pyir chain imports + # from this module at module top). + from .pyir_runtime import _has_decomposable_staged_fields + + if _has_decomposable_staged_fields(old_value): + return None + raise DSLUserCodeError( DiagId.PHASE_MUTATE_PYTHON, filename=filename, @@ -402,12 +523,16 @@ def assign_meta_staged_check( if auto_m2s: import warnings + # The constructor-suggestion prefix is per-DSL (declared on the env + # manager), so base_dsl names no dialect; empty means no namespace. + ns = getattr(env_manager, "dsl_constructor_namespace", "") + ctor_prefix = f"{ns}." if ns else "" warnings.warn( f"Implicit Meta-to-Staged promotion of `{target_name}` " f"(from {type(old_value).__name__} to " f"{type(new_value).__name__}) is deprecated. " f"Initialize as: {target_name} = " - f"cute.{type(new_value).__name__}({old_value!r})\n" + f"{ctor_prefix}{type(new_value).__name__}({old_value!r})\n" f" at {filename}:{lineno}", DeprecationWarning, stacklevel=3, diff --git a/python/CuTeDSL/cutlass/base_dsl/preprocess_mode.py b/python/CuTeDSL/cutlass/base_dsl/preprocess_mode.py index 843c42e864..ccbe528708 100644 --- a/python/CuTeDSL/cutlass/base_dsl/preprocess_mode.py +++ b/python/CuTeDSL/cutlass/base_dsl/preprocess_mode.py @@ -25,19 +25,19 @@ class _PreprocessModeState: def __init__(self, dsl: "BaseDSL") -> None: self._dsl = dsl - self._epoch = 0 self._stack: list[tuple[bool, bool, DSLPreprocessor]] = [] - def stamp(self, preprocessor: DSLPreprocessor) -> DSLPreprocessor: - setattr(preprocessor, "_mode_epoch", self._epoch) - return preprocessor - def current_signature( self, - ) -> tuple[tuple[bool, bool, str | None, int]]: + ) -> tuple[tuple[bool, bool, str | None]]: return (self.current_mode_key(),) - def current_mode_key(self) -> tuple[bool, bool, str | None, int]: + def current_mode_key(self) -> tuple[bool, bool, str | None]: + # The rewrite output is a pure function of (source, mode key): the + # preprocessor class fixes the choke set, and the two flags fix its + # parameters. No history/epoch component — restoring an identical + # mode MUST yield an identical key, or every cached rewrite in the + # process is spuriously invalidated on each mode round trip. preprocessor = getattr(self._dsl, "preprocessor", None) preprocessor_cls = ( type(preprocessor).__name__ if preprocessor is not None else None @@ -46,7 +46,6 @@ def current_mode_key(self) -> tuple[bool, bool, str | None, int]: bool(self._dsl.envar.enable_pyir), bool(self._dsl.envar.auto_m2s), preprocessor_cls, - self._epoch, ) def push_pyir(self) -> None: @@ -89,18 +88,14 @@ def _apply( auto_m2s: bool, preprocessor: DSLPreprocessor, ) -> None: - current_mode = ( - bool(self._dsl.envar.enable_pyir), - bool(self._dsl.envar.auto_m2s), - type(self._dsl.preprocessor).__name__, - ) - next_mode = ( - bool(enable_pyir), - bool(auto_m2s), - type(preprocessor).__name__, - ) - if current_mode != next_mode: - self._epoch += 1 + if enable_pyir: + # Turning PyIR on here bypasses the environment manager, which is + # where _pyir_register_mode_fact would otherwise bind the Numeric + # write funnel. Without it the wrapper write funnel is undeclared + # for this trace and no holder write is ever stamped. + from .typing import _pyir_install_numeric_write_funnel + + _pyir_install_numeric_write_funnel() self._dsl.envar.enable_pyir = enable_pyir self._dsl.envar.auto_m2s = auto_m2s - self._dsl.preprocessor = self.stamp(preprocessor) + self._dsl.preprocessor = preprocessor diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_call_boundary.py b/python/CuTeDSL/cutlass/base_dsl/pyir_call_boundary.py index 29903c61e1..ad054b09d0 100644 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_call_boundary.py +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_call_boundary.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: LicenseRef-NvidiaProprietary # # Use of this software is governed by the terms and conditions of the @@ -10,9 +10,2408 @@ # is strictly prohibited. -"""PyIR runtime -- call-boundary layer; see facade for the public surface.""" +"""Call-boundary effect observer: a plain (un-instrumented) callee's stores on +tracked places are diffed at the call site and replayed through the chokes. -from .pyir_entrypoints import * # noqa: F401,F403 (re-export lower layers up the chain) +DECLARED observation domain of the snapshot/diff walk: objects reachable from +the call's receiver, arguments, kwargs values, closure cells, the callee's +TAKEN parameter defaults (``__defaults__``/``__kwdefaults__`` objects the call +binding leaves unbound, LangRef §8.7), and the registered slot holders; per +object -- its instance storage (``__dict__`` + +set ``__slots__``), raw dict/list legs, closure-cell contents, and the +class-level data attributes of its USER-MODULE class (owner = the class +object). A returned object born in the callee enters the identity domain at +the boundary return (owner-token mint + born-class stamp); its leaves publish +at their binding chokes. Writes outside this domain are detected at their +next choked read or stay on the declared floor.""" +import bisect as _bisect +import builtins as _builtins +import collections as _collections +import dataclasses as _dataclasses +import dis as _dis +import functools as _functools +import heapq as _heapq +import operator as _operator +import os as _os +import sys as _sys +import types as _types +import weakref as _weakref -__all__ = [name for name in list(globals()) if not name.startswith("__")] +from .pyir_state import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_core import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_corewalk import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_loop_carry import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_entrypoints import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_state import _Sentinel + +from . import pyir_class_facts as _pyir_class_facts + + +# -- Callee module fact (fast-path): user code vs DSL/stdlib/third-party. + +# The top-level package this base_dsl instance is vendored under (e.g. +# ``cutlass`` for ``cutlass.base_dsl``); its whole namespace is the DSL. +_PYIR_BOUNDARY_DSL_TOP = __name__.split(".")[0] if "." in __name__ else None + +# The packages root the DSL is installed under (parent of the top package): +# any module whose spec origin lives under it is DSL-distribution code. +_PYIR_BOUNDARY_PKG_ROOTS: "tuple[str, ...]" = () +try: + _pkg_root = _os.path.dirname( + _os.path.dirname(_os.path.dirname(_os.path.abspath(__file__))) + ) + _PYIR_BOUNDARY_PKG_ROOTS = tuple( + {_pkg_root + _os.sep, _os.path.realpath(_pkg_root) + _os.sep} + ) +except Exception: # pragma: no cover - defensive + _PYIR_BOUNDARY_PKG_ROOTS = () + +_PYIR_BOUNDARY_MODULE_VERDICT: "dict[str, bool]" = {} + + +def _pyir_boundary_module_is_user(mod: "str | None") -> bool: + """MODULE fact: does *mod* resolve to user code? DSL packages, declared + library packages, stdlib, and site-packages are not user code (skip).""" + if mod is None: + return False # C extension / builtin: no resolvable origin + if not mod or mod == "__main__": + return True + cached = _PYIR_BOUNDARY_MODULE_VERDICT.get(mod) + if cached is not None: + return cached + verdict = True + top = mod.split(".", 1)[0] + try: + if _PYIR_BOUNDARY_DSL_TOP is not None and top == _PYIR_BOUNDARY_DSL_TOP: + verdict = False + elif any( + mod == p or mod.startswith(p + ".") + for p in _pyir_class_facts._DECLARED_PREFIXES + ): + verdict = False + elif top in getattr(_sys, "stdlib_module_names", ()): + verdict = False + else: + import importlib.util as _importlib_util + + try: + spec = _importlib_util.find_spec(top) + except (ModuleNotFoundError, ValueError, AttributeError, ImportError): + spec = None + origin = getattr(spec, "origin", None) if spec is not None else None + if origin is None: + # Namespace package / frozen / builtin: try the submodule + # search locations before giving up. + locations = ( + getattr(spec, "submodule_search_locations", None) + if spec is not None + else None + ) + if locations: + origin = next(iter(locations), None) + if origin is not None: + if "site-packages" in origin or "dist-packages" in origin: + verdict = False + elif any( + origin.startswith(root) for root in _PYIR_BOUNDARY_PKG_ROOTS + ) or any( + _os.path.realpath(origin).startswith(root) + for root in _PYIR_BOUNDARY_PKG_ROOTS + ): + verdict = False + except Exception: + # Fail-soft: observation of a non-user callee is a no-op diff. + verdict = True + _PYIR_BOUNDARY_MODULE_VERDICT[mod] = verdict + return verdict + + +# Wrapper-consuming layers: the DSL's own Python namespace takes watched-META +# wrappers by design (the retargetable ``ir_value`` channel); the compiled +# ``_mlir`` builder surface, stdlib, and third-party code take plain payloads. +_PYIR_BOUNDARY_WRAPPER_CONSUMER_VERDICT: "dict[str, bool]" = {} + + +def _pyir_boundary_module_consumes_wrappers(mod: "str | None") -> bool: + """MODULE fact: does *mod* consume watched-META wrappers by design?""" + if mod is None or _PYIR_BOUNDARY_DSL_TOP is None: + return False + cached = _PYIR_BOUNDARY_WRAPPER_CONSUMER_VERDICT.get(mod) + if cached is not None: + return cached + top = _PYIR_BOUNDARY_DSL_TOP + verdict = (mod == top or mod.startswith(top + ".")) and not ( + mod == top + "._mlir" or mod.startswith(top + "._mlir.") + ) + _PYIR_BOUNDARY_WRAPPER_CONSUMER_VERDICT[mod] = verdict + return verdict + + +# -- Dispatch: _pyir_call_boundary_(callee) -> callee | proxy + +# Re-entrancy: function ids currently executing under an observation, so a +# recursive self-call is observed once (at the outermost boundary). +_PYIR_BOUNDARY_ACTIVE: "set[int]" = set() + + +class _PyirCallBoundaryProxy: + """Transparent call-through that observes the callee's effects on tracked + places (snapshot -> plain call -> diff -> instrumented replay).""" + + # Slot attribute types (assigned via ``object.__setattr__`` in ``__init__`` + # to stay off the proxy's attribute-forwarding path). + _pyir_boundary_callee: Any + _pyir_boundary_sc_rhs: bool + + __slots__ = ("_pyir_boundary_callee", "_pyir_boundary_sc_rhs") + + def __init__(self, callee: Any, sc_rhs: bool = False) -> None: + _pyir_setattr_raw(self, "_pyir_boundary_callee", callee) + _pyir_setattr_raw(self, "_pyir_boundary_sc_rhs", sc_rhs) + + def __repr__(self) -> str: # pragma: no cover - debug aid + return f"" + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + callee = self._pyir_boundary_callee + try: + active = _pyir_boundary_trace_active() + except Exception: + active = False + if not active: + return callee(*args, **kwargs) + _pyir_global_write_guard(callee) + func = getattr(callee, "__func__", callee) + fid = id(func) + if fid in _PYIR_BOUNDARY_ACTIVE: + return callee(*args, **kwargs) + try: + roots = _pyir_boundary_tracked_roots(callee, args, kwargs) + except Exception: + roots = [] + # F-SPEC: taken parameter defaults are consumed with no read choke; + # seal their payload rows whether or not tracked state is reachable. + try: + _pyir_spec_boundary_default_reads(callee, args, kwargs) + except Exception: + pass + if not roots: + watch = _pyir_global_write_watch(callee) + result = callee(*args, **kwargs) + _pyir_global_write_check(watch) + return result + try: + _pyir_boundary_record_meta_cell_reads(callee) + except Exception: + pass + # Read-half tier 1 (dominance refresh): re-bind an escaped staged attr + # leaf to a fresh dominating load BEFORE the callee reads it raw. + try: + _pyir_reload_stale_staged_attr_leaves("call boundary in") + except Exception: + pass + # Read-half tier 3 (imminent-write pre-stage): fields the callee's own + # code writes run on STAGED values, so the RMW derives from the cell. + try: + prestaged = _pyir_boundary_prestage_imminent_writes(callee, args) + except DSLUserCodeError: + raise # the off-mode promotion refusal must reach the author + except Exception: + prestaged = [] + # Read-half tier 2 (meta read-attribution): plain-meta attr leaves get + # place-attributed wrappers so reads record retargetable constants. + try: + wrapped = _pyir_boundary_stage_meta_reads(roots) + except Exception: + wrapped = [] + watch = _pyir_global_write_watch(callee) + snapshot = _pyir_boundary_snapshot(roots) + try: + _PYIR_BOUNDARY_ACTIVE.add(fid) + _PYIR_BOUNDARY_CALLEE_DEPTH[0] += 1 + try: + result = callee(*args, **kwargs) + finally: + _PYIR_BOUNDARY_CALLEE_DEPTH[0] -= 1 + _PYIR_BOUNDARY_ACTIVE.discard(fid) + # Reached only on a normal return: a raising callee pops the + # snapshot with no emission (the outer finally still restores the + # tier-2/3 binds so no wrapper stays bound in user storage). + _pyir_boundary_commit( + callee, snapshot, self._pyir_boundary_sc_rhs, result=result + ) + finally: + try: + _pyir_boundary_restore_meta_reads(wrapped) + except Exception: + pass + try: + _pyir_boundary_restore_prestage(prestaged) + except Exception: + pass + _pyir_global_write_check(watch) + return result + + +def _pyir_boundary_trace_active() -> bool: + """True while a PyIR-instrumented function body is tracing into an + open MLIR function (the only scope where a boundary commit can emit).""" + if pyir is None or not _PYIR_SCOPE_STACK: + return False + return _get_function_entry_block() is not None + + +# Numeric-protocol builtins (LangRef 3.3.8): with a watched receiver they +# consume through the registered dunder arms, so the funnel must not pre-unwrap. +_PYIR_PROTOCOL_BUILTINS: "tuple[Any, ...]" = (divmod, pow, complex) + + +class _PyirMetaUnwrapProxy: + """Unwrap funnel for callees outside the wrapper-consuming layers: each + watched-META argument is consumed HERE -- recorded and passed as its plain + payload -- so a compiled/stdlib callee never receives a wrapper.""" + + _pyir_unwrap_callee: Any + + __slots__ = ("_pyir_unwrap_callee",) + + def __init__(self, callee: Any) -> None: + self._pyir_unwrap_callee = callee + + def __call__(self, *args: Any, **kwargs: Any) -> Any: + callee = self._pyir_unwrap_callee + try: + active = _pyir_boundary_trace_active() + except Exception: + active = False + if active: + _pyir_global_write_guard(callee) + # A watched RECEIVER guarantees dunder interception (subclass + # priority); a non-watched lhs C-slot needs the funnel witness. + if ( + callee in _PYIR_PROTOCOL_BUILTINS + and args + and isinstance(args[0], _WatchedM) + ): + return callee(*args, **kwargs) + if any(isinstance(a, _WatchedM) for a in args): + args = tuple(_pyir_boundary_consume_meta_arg(a, callee) for a in args) + if kwargs and any(isinstance(v, _WatchedM) for v in kwargs.values()): + kwargs = { + k: _pyir_boundary_consume_meta_arg(v, callee) + for k, v in kwargs.items() + } + watch = _pyir_global_write_watch(callee) + result = callee(*args, **kwargs) + _pyir_global_write_check(watch) + return result + return callee(*args, **kwargs) + + +def _pyir_boundary_consume_meta_arg(value: Any, callee: Any = None) -> Any: + """Record one watched-META consumption at the funnel and return the plain + payload; a promoted place holds no trace-time constant, so it refuses.""" + if not isinstance(value, _WatchedM): + return value + slot_key = value._slot_key + if slot_key is not None and _slot_refs.get(slot_key) is not None: + filename, lineno = _first_non_dsl_caller_location() + callee_name = getattr(callee, "__name__", None) if callee is not None else None + raise DSLUserCodeError( + DiagId.ATTR_BUILDER_REQUIRES_CONSTANT, + filename=filename, + lineno=lineno, + callee=callee_name or "this MLIR builder call", + var=str(slot_key[-1]) if isinstance(slot_key, tuple) else str(slot_key), + ) + value._record_structural_consumption() + return value._pyir_raw_payload + + +# -- Reflection-primitive routing (F-STORE / F-READ arm): Python's own fixed +# attribute vocabulary, keyed by FUNCTION-OBJECT identity, routed to the +# read/write chokes with owner = arg0 and slot = the runtime-exact string. + + +def _pyir_reflect_storage_owner(obj: Any, name: str) -> Any: + """The namespace owner a runtime-named attribute access resolves to: the + instance for instance storage, the defining class for a class-level data + attribute (F-SHAPE), ``None`` for computed attributes (place is ⊥).""" + if obj is None or _is_staged_value(obj) or isinstance(obj, _WatchedM): + return None + if not isinstance(obj, type): + items = _instance_storage_items(obj) + if items is not None and name in items: + return obj + mro = getattr(obj if isinstance(obj, type) else type(obj), "__mro__", ()) + for klass in mro: + if name in klass.__dict__: + if name in _pyir_class_facts.class_storage_items(klass): + return klass + return None # descriptor / dunder / nested class: computed + return None + + +def _pyir_routed_getattr(obj: Any, name: Any, *default: Any) -> Any: + """Routed ``getattr``: the native dereference is Python truth; a storage hit + re-resolves through the read choke (records the read, loads live cells); a + hit that resolves to NO storage slot is judged as a fabricated read; a + tolerated default-taken MISS seals an absence row (the bake depends on the + name staying absent).""" + try: + cur = getattr(obj, name) + except AttributeError: + if not default: + raise + if _pyir_boundary_trace_active() and isinstance(name, str): + _pyir_spec_record_attr_probe(obj, name, present=False) + return default[0] + if not _pyir_boundary_trace_active() or not isinstance(name, str): + return cur + owner = _pyir_reflect_storage_owner(obj, name) + if owner is None: + _pyir_judge_fabricated_attr_read(obj, name, cur) + return cur + return pyir_read(f".{name}", cur, owner=owner, slot_name=name) + + +def _pyir_routed_hasattr(obj: Any, name: Any) -> bool: + """Routed ``hasattr`` (CPython's own algorithm: one ``getattr``, catching + AttributeError): both answers seal a presence/absence probe row (the bake + is the boolean, never the value), and a computed hit is judged like the + getattr leg; the boolean answer is Python's.""" + try: + cur = getattr(obj, name) + except AttributeError: + if _pyir_boundary_trace_active() and isinstance(name, str): + _pyir_spec_record_attr_probe(obj, name, present=False) + return False + if _pyir_boundary_trace_active() and isinstance(name, str): + if _pyir_reflect_storage_owner(obj, name) is None: + _pyir_judge_fabricated_attr_read(obj, name, cur) + # Presence of a region-conditional first-def is path-dependent: same + # read-before-set refusal as a value read. + _pyir_check_cf_attr_first_def_read(obj, name) + _pyir_spec_record_attr_probe(obj, name, present=True) + return True + + +def _pyir_reflect_store(setter: Any, obj: Any, name: Any, value: Any) -> None: + """Routed attribute store: the native write runs first (Python truth), then + a STORAGE write replays through the write choke at this position -- the + same lifecycle as an instrumented ``obj.name = value``.""" + if not _pyir_boundary_trace_active() or not isinstance(name, str): + setter(obj, name, value) + return + pre = _pyir_boundary_current_value(obj, name) + setter(obj, name, value) + if _pyir_boundary_current_value(obj, name) is not value: + return # a descriptor rerouted the store: not a storage write here + label = f".{name}" + filename, lineno = _first_non_dsl_caller_location() + old = ( + None + if pre is _PYIR_BOUNDARY_MISSING + else pyir_read(label, pre, owner=obj, slot_name=name) + ) + result = pyir_assign(label, old, value, filename, lineno, owner=obj, slot_name=name) + if result is not value: + _pyir_holder_store(obj, name, result) + # Meta first-def through reflection: record the binding position (a + # staged first-def records through the assign choke). + if pre is _PYIR_BOUNDARY_MISSING and not _is_staged_value(value): + _pyir_record_cf_attr_first_def(obj, name) + + +def _pyir_routed_setattr(obj: Any, name: Any, value: Any) -> None: + _pyir_reflect_store(setattr, obj, name, value) + + +def _pyir_routed_object_setattr(obj: Any, name: Any, value: Any) -> None: + _pyir_reflect_store(_pyir_setattr_raw, obj, name, value) + + +def _pyir_delete_attr(obj: Any, name: Any) -> None: + """The one attribute-deletion funnel (syntactic ``del obj.name`` and routed + ``delattr``; LangRef 3.12 sections 7.5 and 3.3.2): deletion is an + unstageable binding effect, so inside dynamic staged CF it cannot be + predicated on the region's runtime condition and refuses loudly. In meta + flow the native delete runs first (Python truth, honoring ``__delattr__`` + and descriptors), then the ledger row dies with the binding (W2) and a + rooted place seals trace-exit ABSENCE (F-SPEC).""" + if not _pyir_boundary_trace_active() or not isinstance(name, str): + delattr(obj, name) + return + if is_inside_staged_cf(): + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.UNSUP_DEL_IN_STAGED_CF, + filename=filename, + lineno=lineno, + obj=obj.__name__ if isinstance(obj, type) else type(obj).__name__, + attr=name, + ) + delattr(obj, name) + _pyir_retire_place_row(None, obj, name) + _pyir_spec_record_unbind(obj, name) + + +def _pyir_routed_delattr(obj: Any, name: Any) -> None: + """Routed ``delattr``: the same funnel as an instrumented ``del obj.name`` + (one choke for both spellings, by construction).""" + _pyir_delete_attr(obj, name) + + +def _pyir_import_guard( + stmt: str, module: "str | None", level: int, package: "str | None", *names: str +) -> None: + """Staged-CF wall ahead of an in-body import statement (LangRef 3.12 + section 7.11): a FIRST import executes the module's code once at trace + time, a side effect that cannot be predicated on a staged region's + runtime condition -- inside dynamic staged CF it refuses loudly, before + that code runs. A pure module-cache hit (section 5.3.1) executes no + module code and passes: its bindings are ordinary meta reads, recorded + as spec roots. Meta flow passes through untouched.""" + if not _pyir_boundary_trace_active() or not is_inside_staged_cf(): + return + if _pyir_import_is_cached(module, level, names, package): + return + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.UNSUP_IMPORT_IN_STAGED_CF, + filename=filename, + lineno=lineno, + stmt=stmt, + ) + + +_PYIR_GLOBAL_WRITE_SCANS: "dict[Any, str | None]" = {} + + +def _pyir_callee_global_write(callee: Any) -> "str | None": + """Name of the first module global the callee's OWN bytecode writes + (STORE_GLOBAL/DELETE_GLOBAL; one frame deep by construction), or None. + Only user-module callees report -- a DSL/stdlib global write is trace-time + machinery, not carried program state.""" + func = getattr(callee, "__func__", callee) + code = getattr(func, "__code__", None) + if code is None: + return None + if code in _PYIR_GLOBAL_WRITE_SCANS: + return _PYIR_GLOBAL_WRITE_SCANS[code] + name = None + if _pyir_boundary_module_is_user(getattr(func, "__module__", None)): + for ins in _dis.get_instructions(code): + if ins.opname in ("STORE_GLOBAL", "DELETE_GLOBAL"): + name = str(ins.argval) + break + _PYIR_GLOBAL_WRITE_SCANS[code] = name + return name + + +def _pyir_global_write_guard(callee: Any) -> None: + """Staged-CF wall ahead of a callee that writes a module global: the write + executes once at trace time, not once per runtime iteration, so every + later read of that global is frozen at its first traced value. Constexpr + regions unroll at trace time (the callee really runs per iteration) and + pass; so does straight-line code (one execution, one write).""" + if not is_inside_staged_cf() or is_inside_constexpr_loop(): + return + name = _pyir_callee_global_write(callee) + if name is None: + return + func = getattr(callee, "__func__", callee) + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.UNSUP_GLOBAL_WRITE_IN_STAGED_CF, + filename=filename, + lineno=lineno, + callee=getattr(func, "__qualname__", None) or "the called function", + name=name, + ) + + +def _pyir_global_write_watch(callee: Any) -> "tuple | None": + """Pre-call half of the two-hop backstop: for a user-module callee inside + dynamic staged CF, capture the callee's module dict and its current + bindings, so a transitive global write (invisible to the one-frame + bytecode scan) is caught as a changed binding after the call.""" + if not is_inside_staged_cf() or is_inside_constexpr_loop(): + return None + func = getattr(callee, "__func__", callee) + mod = getattr(func, "__module__", None) + if mod is None or not _pyir_boundary_module_is_user(mod): + return None + d = getattr(_sys.modules.get(mod), "__dict__", None) + if d is None: + return None + before = {k: v for k, v in d.items() if not k.startswith("__")} + return (callee, d, before) + + +def _pyir_global_write_check(watch: "tuple | None") -> None: + """Post-call half of the two-hop backstop: refuse on any rebound, added, + or deleted module-global binding (in-place container mutation does not + rebind and passes).""" + if watch is None: + return + callee, d, before = watch + changed = None + for k, v in d.items(): + if not k.startswith("__") and (k not in before or before[k] is not v): + changed = k + break + if changed is None: + for k in before: + if k not in d: + changed = k + break + if changed is None: + return + func = getattr(callee, "__func__", callee) + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.UNSUP_GLOBAL_WRITE_IN_STAGED_CF, + filename=filename, + lineno=lineno, + callee=getattr(func, "__qualname__", None) or "the called function", + name=changed, + ) + + +def _pyir_import_record( + module: "str | None", level: int, package: "str | None", *pairs: tuple +) -> None: + """Record arm behind an executed in-body import (meta flow): each bound + name becomes a verified spec root (F-SPEC). ``module is None, level 0`` + is the plain-``import`` form; *package* (the defining module's declared + package, a rewrite-time constant) anchors a relative ``from``-import + exactly as the import system resolved it.""" + if not _pyir_boundary_trace_active(): + return + _pyir_spec_record_import(module, level, pairs, package) + + +def _pyir_routed_vars(*args: Any) -> Any: + """Routed ``vars``: each instance-storage leaf read is recorded through the + read choke, and the returned mapping funnels through the same read choke as + an ``obj.__dict__`` attr read (``vars(obj)`` IS ``obj.__dict__`` -- one + adoption funnel, so writes into it carry instead of silently dropping).""" + if not args: + # Zero-arg ``vars()`` is the CALLER's locals; evaluating it here would + # name this frame instead. + return _sys._getframe(1).f_locals + result = vars(*args) + if _pyir_boundary_trace_active() and isinstance(result, dict): + obj = args[0] + if not (_is_staged_value(obj) or isinstance(obj, _WatchedM)): + for k, v in list(result.items()): + if isinstance(k, str) and not k.startswith("__"): + pyir_read(f".{k}", v, owner=obj, slot_name=k) + result = pyir_read("", result, owner=obj, slot_name="__dict__") + return result + + +def _pyir_routed_attrgetter(*names: Any) -> Any: + """Routed ``operator.attrgetter``: the returned getter resolves each hop + through the routed ``getattr`` while a trace is active.""" + native = _operator.attrgetter(*names) # native validation of the names + + def _get_path(obj: Any, dotted: str) -> Any: + for hop in dotted.split("."): + obj = _pyir_routed_getattr(obj, hop) + return obj + + def _routed_getter(obj: Any) -> Any: + if not _pyir_boundary_trace_active(): + return native(obj) + if len(names) == 1: + return _get_path(obj, names[0]) + return tuple(_get_path(obj, n) for n in names) + + return _routed_getter + + +class _PyirSynthLocalsView(dict): + """``locals()`` mapping of a SYNTHESIZED scope (a rewritten if/ifexp/while + arm holds only the names it rebinds, carried as parameters): a HIT is the + live carried value (Python truth); a MISS or an enumeration cannot be + answered from the partial view, so it refuses loudly instead of silently + diverging from the source program's full-scope ``locals()``.""" + + def _refuse(self, name: Any) -> Any: + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.SCOPE_LOCALS_SYNTH_MISS, + filename=filename, + lineno=lineno, + name=str(name), + ) + + def __getitem__(self, key: Any) -> Any: + if dict.__contains__(self, key): + return dict.__getitem__(self, key) + return self._refuse(key) + + def get(self, key: Any, default: Any = None) -> Any: + if dict.__contains__(self, key): + return dict.__getitem__(self, key) + return self._refuse(key) + + def __contains__(self, key: Any) -> bool: + if dict.__contains__(self, key): + return True + return bool(self._refuse(key)) + + def _refuse_enumeration(self) -> Any: + return self._refuse("") + + def __iter__(self) -> Any: + return self._refuse_enumeration() + + def __len__(self) -> int: + return int(self._refuse_enumeration()) + + def keys(self) -> Any: # type: ignore[override] + return self._refuse_enumeration() + + def items(self) -> Any: # type: ignore[override] + return self._refuse_enumeration() + + def values(self) -> Any: # type: ignore[override] + return self._refuse_enumeration() + + def __repr__(self) -> str: + return "" + + +def _pyir_traced_locals(mapping: Any, callee: Any, in_synth: bool) -> Any: + """Choke for a bare zero-arg ``locals()``/``vars()`` in a rewritten body + (library ref section `vars`: zero-arg ``vars()`` == ``locals()``). At + function scope the frame mapping keeps source names live (pass-through); + in a preprocessor-SYNTHESIZED arm scope the mapping is a partial view, so + lookups are guarded (hit = truth, miss/enumeration = loud). A shadowed + builtin name keeps whatever the shadow returned (Python truth).""" + if callee is not _builtins.locals and callee is not _builtins.vars: + return mapping + if not in_synth or not isinstance(mapping, dict): + return mapping + if not _pyir_boundary_trace_active(): + return mapping + return _PyirSynthLocalsView(mapping) + + +def _pyir_routed_dict_view(view: Any) -> Any: + """Routed unbound ``dict.keys/values/items`` (``dict.items(d)``): the + bound spelling routes at the dispatcher's receiver check; the unbound + base-class call carries the receiver as ``args[0]`` and must hit the same + iteration choke before the view raw-iterates storage. Machinery's direct + raw walks never dispatch through the boundary, so they stay unobserved.""" + + def _routed(*args: Any, **kwargs: Any) -> Any: + if ( + args + and isinstance(args[0], _WatchedDict) + and _WATCHED_DICT_ITER_HOOK[0] is not None + ): + served = _WATCHED_DICT_ITER_HOOK[0](args[0], f".{view.__name__}()") + if served is not None: + return served + return view(*args, **kwargs) + + return _routed + + +# The Python-fixed primitive set, keyed by function-object identity (the +# primitives are immortal; the id key can never be recycled under them). +_PYIR_REFLECTION_ROUTES: "dict[int, tuple[Any, Any]]" = { + id(prim): (prim, router) + for prim, router in ( + (getattr, _pyir_routed_getattr), + (setattr, _pyir_routed_setattr), + (delattr, _pyir_routed_delattr), + (hasattr, _pyir_routed_hasattr), + (object.__setattr__, _pyir_routed_object_setattr), + (vars, _pyir_routed_vars), + (_operator.attrgetter, _pyir_routed_attrgetter), + (dict.keys, _pyir_routed_dict_view(dict.keys)), + (dict.values, _pyir_routed_dict_view(dict.values)), + (dict.items, _pyir_routed_dict_view(dict.items)), + ) +} + + +def _pyir_routed_reduce(*args: Any, **kwargs: Any) -> Any: + """Routed ``functools.reduce``: the callee argument crosses a HOF boundary, + so every application dispatches through the standard call boundary.""" + if args: + args = (_pyir_call_boundary_(args[0]),) + args[1:] + elif "function" in kwargs: + kwargs = {**kwargs, "function": _pyir_call_boundary_(kwargs["function"])} + return _functools.reduce(*args, **kwargs) + + +# Stdlib higher-order drivers whose CALLEE argument must cross the boundary, +# keyed by function-object identity exactly like the reflection routes. +_PYIR_HOF_ROUTES: "dict[int, tuple[Any, Any]]" = { + id(_functools.reduce): (_functools.reduce, _pyir_routed_reduce), +} + + +def _pyir_routed_clevel_mutator(fn: Any, spelling: str) -> Any: + """Routed C-implemented stdlib container mutator (the ``heapq`` module + functions / ``bisect.insort`` family): the mutation happens inside C + code, so it performs no Python-level attribute call or subscript store + for the method guard or the subscript-write choke to see -- a silent + trace-time freeze. Refuse when a meta list/deque argument would be + mutated inside dynamic staged CF; otherwise delegate to the standard + unwrap funnel.""" + funnel = _PyirMetaUnwrapProxy(fn) + module_name, _, func_name = spelling.rpartition(".") + + def _routed(*args: Any, **kwargs: Any) -> Any: + if is_inside_staged_cf() and not is_inside_constexpr_loop(): + # Every routed function mutates ONLY its first argument (the + # heap / the sorted sequence; ``bisect.insort`` also accepts it + # as keyword ``a``). An item argument that happens to be a + # list (``heappush(heap, [k, v])``) is stored, never mutated. + # The heapq routes are C-enforced list-only; the insort family + # also mutates any duck-typed sequence through ``.insert``. + mutated = args[0] if args else kwargs.get("a") + if isinstance(mutated, (list, _collections.deque)) or ( + module_name == "bisect" and hasattr(mutated, "insert") + ): + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.UNSUP_META_CONTAINER_MUTATION, + filename=filename, + lineno=lineno, + container=module_name, + method=func_name, + kind=( + "deque" + if isinstance(mutated, _collections.deque) + else "list" + if isinstance(mutated, list) + else type(mutated).__name__ + ), + ) + return funnel(*args, **kwargs) + + return _routed + + +# C-implemented stdlib functions that mutate a caller-owned container +# argument in place, keyed by function-object identity like the reflection +# routes (the table tuple holds the function, so its id cannot be recycled). +# ``bisect.insort`` aliases ``insort_right``; it is listed LAST so the +# diagnostic shows the spelling users overwhelmingly write. The max-heap +# names are version-dependent, hence the getattr walk. +_PYIR_CLEVEL_MUTATOR_ROUTES: "dict[int, tuple[Any, Any]]" = { + id(fn): (fn, _pyir_routed_clevel_mutator(fn, f"{mod_name}.{name}")) + for mod, mod_name, names in ( + ( + _heapq, + "heapq", + ( + "heapify", + "heappush", + "heappop", + "heappushpop", + "heapreplace", + "_heapify_max", + "_heappop_max", + "_heapreplace_max", + "heapify_max", + "heappush_max", + "heappop_max", + "heappushpop_max", + "heapreplace_max", + ), + ), + (_bisect, "bisect", ("insort_left", "insort_right", "insort")), + ) + for name in names + if (fn := getattr(mod, name, None)) is not None +} +# A comprehension's loop variables are comprehension-local, but a walrus +# target binds in the ENCLOSING scope (PEP 572) -- without this, ``fn`` +# leaks into the module namespace and rides the export chain. +del fn + + +# Receiver-carrying spellings of a deque mutation: the method-name guard +# tables key on attribute-call SYNTAX, so an aliased bound method +# (``push = q.append; push(x)``) and the in-place dunders (``q += .../q *= +# ...`` route ``deque.__iadd__``/``__imul__`` through this dispatcher via +# ``_pyir_inplace_binop``) would otherwise mutate silently -- a plain deque +# is never adopted as a watched container, so no object-side choke exists. +_PYIR_DEQUE_BOUND_MUTATORS = _PYIR_DEQUE_MUTATORS | {"__iadd__", "__imul__"} + + +def _pyir_wrapped_user_module(func: Any) -> "str | None": + """WRAPPER fact: the first user module on *func*'s ``__wrapped__`` chain + (the ``functools.update_wrapper`` identity); None when the chain has none.""" + seen: "set[int]" = set() + target = getattr(func, "__wrapped__", None) + while target is not None and id(target) not in seen: + seen.add(id(target)) + mod = getattr(target, "__module__", None) + if _pyir_boundary_module_is_user(mod): + return mod + target = getattr(target, "__wrapped__", None) + return None + + +def _pyir_refuse_memoized_callee(func: Any) -> None: + """Curated wall: a memoized user callee under staged CF skips its body on + a cache hit, so its effects can never replay per runtime execution.""" + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.BOUNDARY_MEMOIZED_IN_STAGED_CF, + filename=filename, + lineno=lineno, + name=getattr(func, "__name__", "this function"), + ) + + +def _pyir_call_boundary_(callee: Any, sc_rhs: bool = False) -> Any: + """Trace-time dispatcher: proxy a plain user-module callable, return anything + else unchanged; *sc_rhs* marks a short-circuited operand (effects refuse).""" + try: + func = getattr(callee, "__func__", callee) + # Memoized wall: an lru_cache hit skips the wrapped USER body, so its + # effects can never replay per runtime execution of a staged region. + if ( + isinstance(func, _functools._lru_cache_wrapper) + and _pyir_boundary_module_is_user( + getattr(getattr(func, "__wrapped__", None), "__module__", None) + ) + and is_inside_staged_cf() + ): + _pyir_refuse_memoized_callee(func) + # Tracked-dict view iteration (``d.keys()/.values()/.items()``): the + # views raw-iterate storage, so pair/value consumption would bypass + # the item read choke -- a trace-time snapshot. User calls all route + # through this dispatcher; inside dynamic staged CF the hook serves + # ``.values()/.items()`` through the per-entry read choke (the + # carried slot reads) and refuses a key set changed inside the region + # (bare ``for k in d`` routes at the object-side ``__iter__``). + _iter_recv = getattr(callee, "__self__", None) + if ( + isinstance(_iter_recv, _WatchedDict) + and getattr(callee, "__name__", None) in ("keys", "values", "items") + and _WATCHED_DICT_ITER_HOOK[0] is not None + ): + _iter_callee = callee + + def _routed_bound_view() -> Any: + served = _WATCHED_DICT_ITER_HOOK[0]( + _iter_recv, f".{_iter_callee.__name__}()" + ) + return _iter_callee() if served is None else served + + return _routed_bound_view + # Deque mutator reached through the callable's own receiver (bound + # method alias, or the in-place dunder ``_pyir_inplace_binop`` + # dispatches here): refuse inside dynamic staged CF exactly like the + # attribute-call guard; trace-time and constexpr mutation stay legal. + if ( + isinstance(_iter_recv, _collections.deque) + and getattr(callee, "__name__", None) in _PYIR_DEQUE_BOUND_MUTATORS + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + ): + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.UNSUP_META_CONTAINER_MUTATION, + filename=filename, + lineno=lineno, + container="deque", + method=callee.__name__, + kind="deque", + ) + # Jit-decorated / DSL-preprocessed callables trace anyway. + if ( + hasattr(func, "_preprocessed") + or hasattr(func, "_dsl_cls") + or hasattr(func, "_dsl_object") + or _pyir_callee_is_rewritten(func) + ): + return callee + if isinstance(callee, (_PyirCallBoundaryProxy, _PyirMetaUnwrapProxy)): + return callee + if not callable(callee): + return callee + # functools.partial: identity and observation belong to ``.func``; the + # frozen args re-enter the dispatched callee as ordinary call args. + # An unchanged ``.func`` (jit / wrapper-consumer fast exit) keeps the + # partial's own stdlib classification below (the unwrap funnel). + if type(callee) is _functools.partial: + inner = _pyir_call_boundary_(callee.func, sc_rhs) + if inner is not callee.func: + return _functools.partial(inner, *callee.args, **callee.keywords) + # Reflection primitives route to the chokes by function-object identity. + route = _PYIR_REFLECTION_ROUTES.get(id(func)) + if route is not None and route[0] is func: + return route[1] + route = _PYIR_HOF_ROUTES.get(id(func)) + if route is not None and route[0] is func: + return route[1] + route = _PYIR_CLEVEL_MUTATOR_ROUTES.get(id(func)) + if route is not None and route[0] is func: + return route[1] + # ``dataclasses.replace`` is a stdlib REBUILD shim: it forwards field + # values untouched into the class's generated ``__init__``, so watched + # meta fields must SURVIVE the call (the ctor-rebind carry keys on + # them). Classify it like the user ctor call it wraps, not like an + # opaque stdlib callee the unwrap funnel de-watches. + if func is _dataclasses.replace: + return _PyirCallBoundaryProxy(callee, sc_rhs) + mod = getattr(func, "__module__", None) + if not _pyir_boundary_module_is_user(mod): + if _pyir_boundary_module_consumes_wrappers(mod): + return callee + # A wrapper masking a user target (``__wrapped__``) is classified + # by the target: its interior user-frame writes must be observed. + wrapped_mod = _pyir_wrapped_user_module(func) + if wrapped_mod is None: + return _PyirMetaUnwrapProxy(callee) + mod = wrapped_mod + # A plain callee's module may carry no jit decoration; funnel it + # through the class-facts materializer so write facts exist by commit. + _pyir_class_facts.ensure_module_materialized(mod) + return _PyirCallBoundaryProxy(callee, sc_rhs) + except DSLUserCodeError: + raise + except Exception as exc: + # Fail-loud wall: a swallowed classification failure would silently + # skip observation, so the callee's effects replay unpredicated. + raise DSLRuntimeError( + "PyIR emission self-check: call-boundary classification failed " + f"for a callee of type {type(callee).__name__}: {exc}" + ) from exc + + +# -- In-place operator routing: an AugAssign dispatches its dunder with no +# ast.Call, so the wrap pass cannot reach it; the preprocessor rewrites the +# statement into this helper, which resolves the in-place slot exactly like +# CPython and routes the bound dunder through the SAME boundary dispatcher. + +_PYIR_INPLACE_OPS: "dict[str, tuple[str, Any, Any]]" = { + "add": ("__iadd__", _operator.add, _operator.iadd), + "sub": ("__isub__", _operator.sub, _operator.isub), + "mul": ("__imul__", _operator.mul, _operator.imul), + "matmul": ("__imatmul__", _operator.matmul, _operator.imatmul), + "truediv": ("__itruediv__", _operator.truediv, _operator.itruediv), + "floordiv": ("__ifloordiv__", _operator.floordiv, _operator.ifloordiv), + "mod": ("__imod__", _operator.mod, _operator.imod), + "pow": ("__ipow__", _operator.pow, _operator.ipow), + "lshift": ("__ilshift__", _operator.lshift, _operator.ilshift), + "rshift": ("__irshift__", _operator.rshift, _operator.irshift), + "or": ("__ior__", _operator.or_, _operator.ior), + "xor": ("__ixor__", _operator.xor, _operator.ixor), + "and": ("__iand__", _operator.and_, _operator.iand), +} + +# Exact types whose in-place slots are C code: the native protocol IS the +# routed protocol (no user frame can hide behind them). +_PYIR_INPLACE_NATIVE_TYPES = frozenset( + ( + int, + float, + bool, + complex, + str, + bytes, + bytearray, + list, + tuple, + dict, + set, + frozenset, + type(None), + ) +) + + +# Property-family attribute stores: the SETTER is class code invoked by the +# interpreter's setattr with no ast.Call; the preprocessor funnels the store +# here so the bound setter routes through the same boundary dispatcher. + +_PYIR_PROPERTY_SETTERS: "dict[tuple[type, str], Any]" = {} + + +def _pyir_property_setter(cls: type, name: str) -> Any: + """The property setter *name* resolves to on *cls* (mro-exact, cached); + ``None`` for ordinary storage attributes and non-property descriptors.""" + key = (cls, name) + if key in _PYIR_PROPERTY_SETTERS: + return _PYIR_PROPERTY_SETTERS[key] + fset = None + for klass in cls.__mro__: + if name in vars(klass): + desc = vars(klass)[name] + if isinstance(desc, property): + fset = desc.fset + break + _PYIR_PROPERTY_SETTERS[key] = fset + return fset + + +_PYIR_SETATTR_OVERRIDES: "dict[type, Any]" = {} + + +def _pyir_user_setattr_override(cls: type) -> Any: + """The user-authored, value-TRANSFORMING ``__setattr__`` governing stores + on *cls* instances (mro-exact, cached); ``None`` when stores are plain or + provably storage-transparent. DSL-internal overrides (struct field + guards, holder write clocks) and transparent user redirects stay on the + rewritten-choke path -- its re-binds replay the override, which is only + safe when the override stores the unmodified value.""" + if cls in _PYIR_SETATTR_OVERRIDES: + return _PYIR_SETATTR_OVERRIDES[cls] + override = None + for klass in cls.__mro__: + if klass is object: + break + fn = vars(klass).get("__setattr__") + if fn is not None: + if _pyir_boundary_module_is_user( + getattr(fn, "__module__", None) + ) and not _pyir_class_facts.setattr_storage_transparent(fn): + override = fn + break + _PYIR_SETATTR_OVERRIDES[cls] = override + return override + + +def _pyir_is_property_store(obj: Any, name: str) -> bool: + """Pure target fact: does storing ``obj.name`` invoke user-defined store + code (a property setter or a user ``__setattr__`` override)? Decides the + store branch BEFORE the RHS evaluates (native order keeps the RHS after + the tracer's pre-store refresh).""" + return ( + _pyir_user_setattr_override(type(obj)) is not None + or _pyir_property_setter(type(obj), name) is not None + ) + + +def _pyir_property_store(obj: Any, name: str, value: Any) -> None: + """Route a store with user-defined store code through it exactly once at + the call boundary (observed, predicated commit of the storage writes). + A ``__setattr__`` override wins over a property, as in CPython; a + setter-less property raises the native AttributeError.""" + override = _pyir_user_setattr_override(type(obj)) + if override is not None: + _pyir_call_boundary_(override.__get__(obj, type(obj)))(name, value) + return + fset = _pyir_property_setter(type(obj), name) + if fset is None: + setattr(obj, name, value) + return + _pyir_call_boundary_(fset.__get__(obj, type(obj)))(value) + + +def _pyir_inplace_binop(op: str, lhs: Any, rhs: Any) -> Any: + """``lhs op= rhs`` VALUE semantics (CPython's in-place protocol) with the + resolved dunder routed through the call boundary, so a user-class + ``__iadd__`` mutating its receiver is observed like any boundary call.""" + dunder, binop, native = _PYIR_INPLACE_OPS[op] + if type(lhs) in _PYIR_INPLACE_NATIVE_TYPES: + return native(lhs, rhs) + # _PyType_Lookup semantics: the slot lives on the type's mro, never on the + # instance, and type-level ``__getattr__`` does not participate. + slot = None + for klass in type(lhs).__mro__: + if dunder in vars(klass): + slot = vars(klass)[dunder] + break + if slot is None: + return binop(lhs, rhs) + if hasattr(type(slot), "__get__"): + bound = slot.__get__(lhs, type(lhs)) + result = _pyir_call_boundary_(bound)(rhs) + else: + result = slot(lhs, rhs) + if result is NotImplemented: + return binop(lhs, rhs) + return result + + +# -- ARGUMENT fact: does the call receive / close over tracked state? + +_PYIR_BOUNDARY_PASSTHROUGH_TYPES = ( + int, + float, + bool, + complex, + str, + bytes, + type(None), +) + + +def _pyir_boundary_value_tracked( + obj: Any, _known_untracked: "set[int] | None" = None +) -> bool: + """Registry lookup: is *obj* state the tracer has already modeled (slot or + candidate holder, adopted container, slot-backed wrapper, or element)? + The gate covers the snapshot walk's whole reachability domain — + containers, instance storage, user-module class objects — iteratively and + unbounded (a bounded proxy would silently skip observation). + + *_known_untracked* shares work across the probes of one root-collection + pass: a probe that returns False explored its whole reachable set, so its + visited ids provably reach no tracked state and later probes skip them. + Only completed-False knowledge is shared -- an early-True probe abandons + its stack, so its visited set proves nothing and is discarded.""" + stack: "list[Any]" = [obj] + visited: "set[int]" = set() + while stack: + cur = stack.pop() + if cur is None or type(cur) in _PYIR_BOUNDARY_PASSTHROUGH_TYPES: + continue + if isinstance(cur, (_WatchedDict, _WatchedList, _WatchedM)): + return True + oid = id(cur) + if oid in _PYIR_SLOT_HOLDERS or oid in _PYIR_CANDIDATE_HOLDERS: + return True + if oid in visited or (_known_untracked is not None and oid in _known_untracked): + continue + visited.add(oid) + if isinstance(cur, (ir.Value, _types.ModuleType)): + continue + if isinstance(cur, type): + # A user-module class object is observable state in its own + # right: its class-level data attributes are boundary places + # (classmethod receivers and class arguments root the diff). + if _pyir_boundary_module_is_user(getattr(cur, "__module__", None)): + return True + continue + if isinstance(cur, dict): + stack.extend(list(dict.values(cur))) + continue + if isinstance(cur, (list, tuple)): + stack.extend(list(cur)) + continue + try: + if getattr(cur, "_mutable_ref", None) is not None: + return True + except Exception: + continue + items = _instance_storage_items(cur) + if items: + stack.extend(list(items.values())) + if _known_untracked is not None: + _known_untracked.update(visited) + return False + + +def _pyir_boundary_default_walk_root(d: Any) -> bool: + """Mutable-shaped default objects the snapshot observes directly: a plain + dict/list (or a tuple reaching one) and user-module-class instances; + scalar, staged, watched, class, and module shapes stay gate-filtered.""" + if d is None or type(d) in _PYIR_BOUNDARY_PASSTHROUGH_TYPES: + return False + if isinstance(d, (_WatchedDict, _WatchedList, _WatchedM, ir.Value)): + return False + if isinstance(d, (dict, list)): + return True + if isinstance(d, tuple): + return any(_pyir_boundary_default_walk_root(v) for v in d) + if isinstance(d, (type, _types.ModuleType)) or _is_staged_value(d): + return False + return ( + _pyir_boundary_module_is_user(getattr(type(d), "__module__", None)) + and _instance_storage_items(d) is not None + ) + + +def _pyir_boundary_tracked_roots( + callee: Any, args: "tuple", kwargs: "dict" +) -> "list[Any]": + """Roots reachable from the call: receiver, args, kwargs values, closure + cell contents, and the callee's taken parameter defaults -- filtered by + the tracked-state registry fact (mutable-shaped taken defaults join + unconditionally: they are only reachable through the callee).""" + roots: "list[Any]" = [] + recv = getattr(callee, "__self__", None) + candidates: "list[Any]" = [] + if recv is not None: + candidates.append(recv) + elif not isinstance( + callee, (_types.FunctionType, _types.BuiltinFunctionType, type) + ): + # A callable-instance callee dispatches through ``type(callee).__call__`` + # with the instance as receiver: its storage is reachable call state. + candidates.append(callee) + candidates.extend(args) + candidates.extend(kwargs.values()) + func = getattr(callee, "__func__", callee) + closure = getattr(func, "__closure__", None) + if closure: + freevars = getattr(getattr(func, "__code__", None), "co_freevars", ()) + for _ci, cell in enumerate(closure): + # The CELL itself is always a root: a ``nonlocal`` rebind lands on + # the cell binding, the holder of its one ``cell_contents`` leg. + if _ci < len(freevars): + # Diagnostics-only: name the closure variable in refusals. + _PYIR_BOUNDARY_CELL_NAMES[_owner_token(cell)] = freevars[_ci] + roots.append(cell) + try: + candidates.append(cell.cell_contents) + except ValueError: + continue + # Taken parameter defaults are persistent def-time state the callee binds + # with no choke (LangRef §8.7): part of the declared observation domain. + for _sel, _dflt in _pyir_boundary_taken_defaults(callee, args, kwargs): + if _pyir_boundary_default_walk_root(_dflt): + roots.append(_dflt) + else: + candidates.append(_dflt) + known_untracked: "set[int]" = set() + for cand in candidates: + if _pyir_boundary_value_tracked(cand, known_untracked): + roots.append(cand) + return roots + + +class _PyirCtxBoundaryProxy: + """WITH-protocol twin of :class:`_PyirCallBoundaryProxy`: observes each dunder + like a boundary call, with the manager AND its class storage as roots.""" + + __slots__ = ("_pyir_ctx_obj",) + + def __init__(self, obj: Any) -> None: + _pyir_setattr_raw(self, "_pyir_ctx_obj", obj) + + def _pyir_observe(self, bound: Any, *args: Any) -> Any: + obj = object.__getattribute__(self, "_pyir_ctx_obj") + try: + active = _pyir_boundary_trace_active() + except Exception: + active = False + if not active: + return bound(*args) + # Same callee fact as the call boundary: a jit-decorated dunder traces + # its own effects, so observing it would double-handle the callee. + _fn = getattr(bound, "__func__", bound) + if ( + hasattr(_fn, "_preprocessed") + or hasattr(_fn, "_dsl_cls") + or hasattr(_fn, "_dsl_object") + or _pyir_callee_is_rewritten(_fn) + ): + return bound(*args) + try: + roots = _pyir_boundary_tracked_roots(bound, args, {}) + except Exception: + roots = [] + roots.append(obj) + _pyir_global_write_guard(bound) + watch = _pyir_global_write_watch(bound) + # The manager's class storage is in the snapshot's declared domain + # (the ctx dispatcher already requires a user-module class). + snapshot = _pyir_boundary_snapshot(roots) + # The WITH protocol is a call boundary: tier 3 pre-stages the fields the + # dunder writes, tier 2 records its meta reads, commit = standard replay. + try: + _pyir_reload_stale_staged_attr_leaves("ctx boundary in") + except Exception: + pass + try: + prestaged = _pyir_boundary_prestage_imminent_writes(bound, args) + except DSLUserCodeError: + raise # the off-mode promotion refusal must reach the author + except Exception: + prestaged = [] + try: + wrapped = _pyir_boundary_stage_meta_reads(roots) + except Exception: + wrapped = [] + try: + _PYIR_BOUNDARY_CALLEE_DEPTH[0] += 1 + try: + result = bound(*args) + finally: + _PYIR_BOUNDARY_CALLEE_DEPTH[0] -= 1 + _pyir_boundary_commit(bound, snapshot, result=result) + finally: + try: + _pyir_boundary_restore_meta_reads(wrapped) + except Exception: + pass + try: + _pyir_boundary_restore_prestage(prestaged) + except Exception: + pass + _pyir_global_write_check(watch) + return result + + def __enter__(self) -> Any: + obj = object.__getattribute__(self, "_pyir_ctx_obj") + return self._pyir_observe(obj.__enter__) + + def __exit__(self, exc_type: Any, exc: Any, tb: Any) -> Any: + obj = object.__getattribute__(self, "_pyir_ctx_obj") + return self._pyir_observe(obj.__exit__, exc_type, exc, tb) + + +def _pyir_ctx_boundary(obj: Any) -> Any: + """Trace-time dispatcher for the WITH-protocol boundary; a dispatch failure + falls back to the unwrapped manager (see :class:`_PyirCtxBoundaryProxy`).""" + try: + if not _pyir_boundary_trace_active(): + return obj + if isinstance(obj, _PyirCtxBoundaryProxy): + return obj + cls = type(obj) + # DSL-authored managers trace their own effects. + if hasattr(cls, "_dsl_cls") or hasattr(cls, "_dsl_object"): + return obj + if not _pyir_boundary_module_is_user(getattr(cls, "__module__", None)): + return obj + if not (hasattr(cls, "__enter__") and hasattr(cls, "__exit__")): + return obj + return _PyirCtxBoundaryProxy(obj) + except Exception: + return obj + + +def _pyir_boundary_record_meta_cell_reads(callee: Any) -> None: + """Record each plain-meta closure cell of *callee* by the CELL's owner + token: a later staged write of a local backed by the same cell (matched + by cell identity, never by name) then refuses loudly at promotion.""" + if not is_inside_staged_cf(): + return + func = getattr(callee, "__func__", callee) + closure = getattr(func, "__closure__", None) + if not closure: + return + code = getattr(func, "__code__", None) + freevars = getattr(code, "co_freevars", ()) + for _ci, cell in enumerate(closure): + try: + contents = cell.cell_contents + except ValueError: + continue + # A scalar shape whose closure read is invisible AND stale-able: a plain + # primitive/_WatchedM, or a staged wrapper reading the pre-region SSA. + _is_scalar_shape = type(contents) in (bool, int, float) or isinstance( + contents, _WatchedM + ) + if not _is_scalar_shape: + try: + _is_scalar_shape = ( + _is_staged_value(contents) + and _can_carry_leaf_ref(contents) + and not _is_compound_single_leaf(contents) + ) + except Exception: + _is_scalar_shape = False + if not _is_scalar_shape: + continue + name = freevars[_ci] if _ci < len(freevars) else "cell_contents" + # F-SPEC: a closure-cell scalar read is a bake of external state; root + # it (or mark the record incomplete when no root is derivable). + _pyir_spec_boundary_closure_read(func, _ci, contents) + tok = _owner_token(cell) + if tok is not None: + _PYIR_BOUNDARY_META_CELL_READS.setdefault( + ("cell", tok), + ( + name, + getattr(code, "co_filename", ""), + getattr(code, "co_firstlineno", 0), + ), + ) + + +# -- Read-half tier 3: imminent-write pre-stage (evidence = the callee's code). + + +def _pyir_ref_use_counts(ref: "ir.Value") -> "tuple[int, int]": + """``(load_count, store_count)`` over *ref*'s current def-use edges.""" + loads = stores = 0 + try: + for use in ref.uses: + user = getattr(use, "owner", None) + user_op = getattr(user, "operation", user) + name = str(getattr(user_op, "name", "")) + if name == "pyir.load": + loads += 1 + elif name == "pyir.store": + stores += 1 + except Exception: + pass + return loads, stores + + +def _pyir_boundary_prestage_imminent_writes( + callee: Any, args: "tuple[Any, ...]" +) -> "list[tuple[Any, str, Any, Any, Any, Any, int, int, bool]]": + """Pre-stage each meta receiver field in the callee's transitive write-facts + (mint cell + fresh load); literal-write facts and celled places skip.""" + if not is_inside_staged_cf(): + return [] + func = getattr(callee, "__func__", callee) + receiver = getattr(callee, "__self__", None) + if receiver is None and args: + receiver = args[0] + if receiver is None or _is_staged_value(receiver): + return [] + # V-7 at the class-facts lookup: facts resolve through type(receiver), so a + # reclassed CELLED receiver would bind the wrong class's write model. + if _pyir_owner_is_celled(receiver): + _pyir_validate_owner_class(receiver) + try: + writes, _complete = _pyir_class_facts.transitive_write_facts(callee) + except Exception: + return [] + if not writes: + return [] + items_map = _instance_storage_items(receiver) + if items_map is None: + return [] + records: "list[tuple[Any, str, Any, Any, Any, Any, int, int, bool]]" = [] + for name in writes: + if name not in items_map: + continue + v = items_map[name] + py_v = v.python_value if isinstance(v, _WatchedM) else v + # A cell-less literal-backed staged scalar field folds at trace time + # exactly like a meta field, so the callee's raw update would chain + # from the entry constant: mint its cell the same way. AUTO_M2S + # only; the default mode keeps the loud boundary refusal. + staged_literal = ( + is_auto_m2s_enabled() + and not isinstance(v, _WatchedM) + and _is_staged_value(v) + and _is_literal_backed(v) + and getattr(v, "_mutable_ref", None) is None + and not _pyir_enclosing_while_cond_is_baked() + ) + if staged_literal: + py_v = v.value + if type(py_v) not in (bool, int, float): + # A compound field whose leaves the region-entry adoption staged: + # the callee reads it RAW, so its rebuild would chain from the + # entry constants, not from the cells the replay stores through. + # Read it through the choke here (same label/owner as the replay, + # so the cells are shared) and bind the load-backed result. + if ( + is_auto_m2s_enabled() + and not isinstance(v, _WatchedM) + and not _is_staged_value(v) + and ( + # The staged-content walk has no dict branch: check the + # entries directly for a dict-valued field. + any(_is_staged_value(e) for e in dict.values(v)) + if isinstance(v, dict) + else _has_any_staged_content(v) + ) + ): + label = _pyir_boundary_label(receiver, name) + try: + if isinstance(v, dict): + # Dict reads do not recurse: read each staged entry + # under the SAME subscript place the replay's + # decomposition stores through, and rebind in place. + for k in list(dict.keys(v)): + e = dict.__getitem__(v, k) + if not _is_staged_value(e): + continue + loaded = pyir_read( + f"{label}[{k!r}]", e, owner=v, slot_name=k + ) + if loaded is not None: + dict.__setitem__(v, k, loaded) + else: + staged_field = pyir_read( + label, v, owner=receiver, slot_name=name + ) + if staged_field is not None: + _pyir_boundary_bind_storage( + receiver, name, staged_field + ) + except Exception: + pass + continue + if _pyir_class_facts.literal_attr_write_fact(func, name) is not None: + continue # latch tier: the loop-carry flip owns literal stores + try: + place = _make_slot_key(None, receiver, name) + except Exception: + place = None + if place is None or _slot_refs.get(place) is not None: + continue + had_recorded_uses = bool(_meta_uses.get(place)) + try: + ref = _meta_promote_slot( + place, + py_v, + display_name=f"{type(receiver).__name__}.{name}", + # Type the ref pointee by the field's own staged type (an + # Int64 field must not mint an Int32 cell). + promoted_value=v if staged_literal else None, + ) + except DSLUserCodeError: + raise # the off-mode promotion refusal must reach the author + except Exception: + continue + if ref is None: + continue + staged = _load_as_dsl(ref, place=place) + if staged is None or not _pyir_boundary_bind_storage(receiver, name, staged): + # Could not bind: retire the fresh row (no store can reach it). + if not had_recorded_uses: + _slot_refs.pop(place, None) + _slot_templates.pop(place, None) + continue + loads0, stores0 = _pyir_ref_use_counts(ref) + records.append( + (receiver, name, staged, v, place, ref, loads0, stores0, had_recorded_uses) + ) + return records + + +def _pyir_boundary_restore_prestage( + records: "list[tuple[Any, str, Any, Any, Any, Any, int, int, bool]]", +) -> None: + """Restore leg: an unfired pre-stage (cell neither stored nor load-consumed) + reverts with no phase residue; a consumed load stays (a real read).""" + for ( + receiver, + name, + staged, + original, + place, + ref, + loads0, + stores0, + had_recorded_uses, + ) in records: + try: + loads1, stores1 = _pyir_ref_use_counts(ref) + if stores1 > stores0 or had_recorded_uses: + continue # written / retargeted-at-mint: the cell is live + consumed = loads1 > loads0 + if not consumed: + raw = _raw_backing_ir_value(staged) + if raw is not None: + try: + consumed = any(True for _ in raw.uses) + except Exception: + consumed = True # cannot prove unused -> keep staged + if consumed: + continue + items_map = _instance_storage_items(receiver) + if items_map is not None and items_map.get(name) is staged: + if _pyir_boundary_bind_storage(receiver, name, original): + _slot_refs.pop(place, None) + _slot_templates.pop(place, None) + except Exception: + continue + + +# -- Read-half tier 2: meta read-attribution for the callee's duration. + + +def _pyir_boundary_bind_storage(holder: Any, name: str, value: Any) -> bool: + """Bind *value* into *holder*'s INSTANCE STORAGE slot *name*, never through a + data descriptor's ``__set__`` (a storage rebind must not run a setter).""" + d = _safe_instance_dict(holder) + if d is not None and name in d: + d[name] = value + return True + for slot_name, desc in _slots_member_descriptors(type(holder)): + if slot_name == name: + desc.__set__(holder, value) + return True + return False + + +def _pyir_boundary_stage_meta_reads( + roots: "list[Any]", +) -> "list[tuple[Any, str, Any, Any]]": + """Bind reachable plain-meta attr leaves to place-attributed ``_WatchedM`` for + the callee's duration, so reads are RECORDED and promotions retarget them.""" + if not is_inside_staged_cf(): + return [] + records: "list[tuple[Any, str, Any, Any]]" = [] + visited: "set[int]" = set() + + def _walk(obj: Any) -> None: + if obj is None or type(obj) in _PYIR_BOUNDARY_PASSTHROUGH_TYPES: + return + if isinstance(obj, (ir.Value, type, _types.ModuleType, _WatchedM)): + return + oid = id(obj) + if oid in visited: + return + visited.add(oid) + if isinstance(obj, _types.CellType): + try: + _walk(obj.cell_contents) + except ValueError: + pass + return + if isinstance(obj, (dict, list)): + vals = list(dict.values(obj)) if isinstance(obj, dict) else list(obj) + for v in vals: + _walk(v) + return + if isinstance(obj, tuple): + for v in obj: + _walk(v) + return + if _is_staged_value(obj): + return + items_map = _instance_storage_items(obj) + if items_map is None: + return + for name, v in list(items_map.items()): + if isinstance(name, str) and name.startswith("__"): + continue + if type(v) in (bool, int, float) and isinstance(name, str): + try: + place = _make_slot_key(None, obj, name) + except Exception: + place = None + if place is None: + continue + wrapper = _WatchedM(v, place) + if _pyir_boundary_bind_storage(obj, name, wrapper): + records.append((obj, name, wrapper, v)) + else: + _walk(v) + + for r in roots: + _walk(r) + return records + + +def _pyir_boundary_stage_taken_default_reads( + callee: Any, args: Any, kwargs: Any +) -> "list[tuple[Any, str, Any, Any]]": + """Tier-2 read binds over an INLINE rewritten callee's taken defaults: a + read-only attr chain in a rewritten body reaches no choke, so plain-meta + leaves bind to place-attributed wrappers exactly as the dispatcher proxy + stages its walk roots. The caller restores the returned records.""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return [] + roots = [ + d + for _sel, d in _pyir_boundary_taken_defaults(callee, args, kwargs) + if not _is_staged_value(d) + ] + return _pyir_boundary_stage_meta_reads(roots) if roots else [] + except Exception: + return [] + + +def _pyir_boundary_restore_meta_reads( + records: "list[tuple[Any, str, Any, Any]]", +) -> None: + """Restore tier-2 binds: a slot still holding this call's wrapper reverts + (recorded uses persist by PLACE); a slot the callee rebound is left.""" + for holder, name, wrapper, original in records: + try: + items_map = _instance_storage_items(holder) + if items_map is not None and items_map.get(name) is wrapper: + _pyir_boundary_bind_storage(holder, name, original) + except Exception: + continue + + +# -- Snapshot: tracked-place leaves reachable from the roots + slot holders. + +_PYIR_BOUNDARY_MISSING = _Sentinel("boundary leaf missing") + +# Diagnostics-only: cell owner token -> closure variable name (tokens are +# identity-true and never recycle, see IdentityKeyedWeakTable). +_PYIR_BOUNDARY_CELL_NAMES: "dict[int | None, str]" = {} + + +def _pyir_boundary_snapshot( + roots: "list[Any]", +) -> "tuple[list[tuple[Any, Any, Any]], dict[int, tuple[Any, frozenset]], set[int]]": + """Record ``(holder, key, pre_value)`` per leaf slot reachable from *roots* + plus registered holders; pre key-sets let added slots replay as first defs; + the visited-id set is the birth criterion for returned objects (see the + module docstring for the declared observation domain).""" + leaves: "list[tuple[Any, Any, Any]]" = [] + key_sets: "dict[int, tuple[Any, frozenset]]" = {} + visited: "set[int]" = set() + + def _walk_class_storage(cls: type) -> None: + # Class-level data attributes of a USER-MODULE class are places owned + # by the CLASS object (F-SHAPE): shared writes diff on the class. + # The user-module portion of the MRO is covered too — an inherited + # data attribute is the same place the reflection read arm attributes + # to its DEFINING class. + for klass in getattr(cls, "__mro__", (cls,)): + cid = id(klass) + if cid in visited: + continue + visited.add(cid) + if not _pyir_boundary_module_is_user(getattr(klass, "__module__", None)): + continue + items = _pyir_class_facts.class_storage_items(klass) + key_sets[cid] = (klass, frozenset(items.keys())) + for name, v in items.items(): + leaves.append((klass, name, v)) + _walk(v) + + def _walk(obj: Any) -> None: + if obj is None or type(obj) in _PYIR_BOUNDARY_PASSTHROUGH_TYPES: + return + if isinstance(obj, type): + # A class reachable as a root/receiver/leaf is observed through + # its class storage (classmethod receivers included). + _walk_class_storage(obj) + return + if isinstance(obj, (ir.Value, _types.ModuleType)): + return + oid = id(obj) + if oid in visited: + return + visited.add(oid) + if isinstance(obj, _WatchedM): + return + if isinstance(obj, _types.CellType): + # Closure cell: one ``cell_contents`` leg; the key is the closure + # variable name when known (diagnostics only). + try: + contents = obj.cell_contents + except ValueError: + return # empty cell -- a later fill is not a place write + leaves.append( + ( + obj, + _PYIR_BOUNDARY_CELL_NAMES.get( + _pyir_lookup_owner_token(obj), "cell_contents" + ), + contents, + ) + ) + _walk(contents) + return + if isinstance(obj, dict): + # An ADOPTED dict already instruments its writes object-level, so its + # legs need no boundary diff; only the VALUES are walked. + watched = isinstance(obj, _WatchedDict) + for k, v in list(dict.items(obj)): + if not watched: + leaves.append((obj, k, v)) + _walk(v) + if not watched: + key_sets[oid] = (obj, frozenset(dict.keys(obj))) + return + if isinstance(obj, list): + watched = isinstance(obj, _WatchedList) + elems = list(obj) + for i, v in enumerate(elems): + if not watched: + leaves.append((obj, i, v)) + _walk(v) + if not watched: + key_sets[oid] = (obj, frozenset(range(len(elems)))) + return + if isinstance(obj, tuple): + for v in obj: + _walk(v) + return + if _is_staged_value(obj): + return + _walk_class_storage(type(obj)) + items_map = _instance_storage_items(obj) + if items_map is None: + return + key_sets[oid] = (obj, frozenset(items_map.keys())) + for name, v in items_map.items(): + if isinstance(name, str) and name.startswith("__"): + continue + leaves.append((obj, name, v)) + _walk(v) + + for r in roots: + _walk(r) + # Registered slot holders: a callee can mutate a holder not passed as an + # argument (module-global, sibling); the registry union keeps it observable. + for entry in list(_PYIR_SLOT_HOLDERS.values()): + holder = entry() if isinstance(entry, _weakref.ref) else entry + if holder is not None: + _walk(holder) + return leaves, key_sets, visited + + +# -- Diff + instrumented replay. + + +def _pyir_boundary_leaf_changed(holder: Any, key: Any, pre: Any, cur: Any) -> bool: + """Did the callee CHANGE this leaf? Exempt: identity, equal meta values, the + same backing SSA, and values backed by the slot's own place cell.""" + if cur is pre: + return False + pre_p = _pyir_unwrap_meta_primitive(pre) + cur_p = _pyir_unwrap_meta_primitive(cur) + if pre_p is not None and cur_p is not None: + try: + return type(pre_p) is not type(cur_p) or pre_p != cur_p + except Exception: + return True + try: + pre_raw = _raw_backing_ir_value(pre) + cur_raw = _raw_backing_ir_value(cur) + except Exception: + return True + if pre_raw is not None and cur_raw is not None and _same_ir_value(pre_raw, cur_raw): + return False + # Same-ref reload exemption: the current value is backed by the slot's own + # cell, so any advance already went through the cell. + try: + mv = _get_slot_mv(holder, key) + except Exception: + mv = None + if mv is not None and mv.ref is not None and _pyir_same_ref_reload(cur, mv.ref): + return False + pre_mv = getattr(pre, "_mutable_ref", None) + if ( + pre_mv is not None + and getattr(pre_mv, "ref", None) is not None + and _pyir_same_ref_reload(cur, pre_mv.ref) + ): + return False + return True + + +def _pyir_boundary_label(holder: Any, key: Any) -> str: + """A diagnostics label for the replayed slot: the ```` prefix cannot + collide with a real access path; the tail names the slot.""" + if isinstance(holder, _types.CellType): + return "" + if isinstance(holder, (dict, list)): + return f"[{key!r}]" + return f".{key}" + + +def _pyir_boundary_current_value(holder: Any, key: Any) -> Any: + """Raw post-call value of the leaf slot (no read chokes fire).""" + if isinstance(holder, type): + return holder.__dict__.get(key, _PYIR_BOUNDARY_MISSING) + if isinstance(holder, _types.CellType): + try: + return holder.cell_contents + except ValueError: + return _PYIR_BOUNDARY_MISSING + if isinstance(holder, dict): + if not dict.__contains__(holder, key): + return _PYIR_BOUNDARY_MISSING + return dict.__getitem__(holder, key) + if isinstance(holder, list): + if not isinstance(key, int) or key >= list.__len__(holder): + return _PYIR_BOUNDARY_MISSING + return list.__getitem__(holder, key) + items_map = _instance_storage_items(holder) + if items_map is None or key not in items_map: + return _PYIR_BOUNDARY_MISSING + return items_map[key] + + +def _pyir_boundary_sc_refuse(key: Any, filename: str, lineno: int) -> None: + """Curated refusal for an effect in a short-circuited ``and``/``or`` operand: + the trace runs it once, so the effect cannot follow the runtime predicate.""" + raise DSLUserCodeError( + DiagId.BOUNDARY_SHORT_CIRCUIT_EFFECT, + name=str(key), + filename=filename, + lineno=lineno, + suggestion=( + f"Restructure the expression as an `if` statement so the update " + f"of `{key}` is a conditional block, or make the guard a " + f"compile-time constant (`const_expr`)." + ), + ) + + +def _pyir_enclosing_while_cond_is_baked() -> bool: + """True when an enclosing ``scf.while``'s already-traced condition is a + baked constant: no cell minted from inside the body can reach it, so a + carried update could never terminate the loop at runtime. Callers keep + the loud refusal instead of committing a carry that would hang.""" + try: + block = ir.InsertionPoint.current.block + except Exception: + return False + for _ in range(256): + if block is None: + return False + try: + parent = block.owner + except Exception: + return False + op = getattr(parent, "operation", parent) + try: + op_name = str(op.name) + except Exception: + return False + if op_name == "scf.while": + try: + before = op.regions[0].blocks[0] + term = before.operations[len(before.operations) - 1] + if str(term.operation.name) == "scf.condition": + cond_owner = term.operands[0].owner + cond_op = getattr(cond_owner, "operation", cond_owner) + if str(cond_op.name) == "arith.constant": + return True + except Exception: + pass + elif _is_func_boundary_op(op_name) or _is_module_boundary_op(op_name): + return False + try: + block = op.block + except Exception: + return False + return False + + +def _pyir_boundary_loop_carry_flip( + callee: Any, + holder: Any, + key: Any, + pre: Any, + pre_p: Any, + cur: Any, + cur_p: Any, + label: str, + filename: str, + lineno: int, +) -> bool: + """Commit a callee's LITERAL store to a loop-carried meta as a carried cell + when its write fact proves iteration-stable guards; False keeps the refusal.""" + if not isinstance(key, str) or cur_p is None: + return False + func = getattr(callee, "__func__", callee) + fact = _pyir_class_facts.literal_attr_write_fact(func, key) + if fact is None: + return False + literals, guard_paths = fact + if not any(type(lit) is type(cur_p) and lit == cur_p for lit in literals): + return False # the facts do not explain the observed value + slot_key = _make_slot_key(None, holder, key) + if slot_key is None: + return False + if _pyir_enclosing_while_cond_is_baked(): + return False # the folded condition could never observe the carry + recorded: "list[tuple[Any, tuple]]" = [] + if guard_paths: + receiver = getattr(callee, "__self__", None) + if receiver is None: + return False # no receiver to resolve the guard places on + single_literal = all( + type(lit) is type(cur_p) and lit == cur_p for lit in literals + ) + for path in guard_paths: + owner: Any = receiver + for i, hop in enumerate(path): + hop_key = _make_slot_key(None, owner, hop) + if hop_key is not None and hop_key == slot_key: + # Self-gated write: re-asserting ONE literal is faithful; + # a multi-literal toggle is not reconstructible. + if not single_literal: + return False + break + try: + value = getattr(owner, hop) + except AttributeError: + return False + raw = _pyir_unwrap_meta_primitive(value) + if raw is None: + if i == len(path) - 1 or _is_staged_value(value): + # The guarding LEAF must be a trace-time meta primitive; + # a staged value at any hop is not provably stable. + return False + if hop_key is None or hop_key in _slot_refs: + return False # guard place already owns a cell + recorded.append((hop_key, (hop, str(key), filename, lineno))) + owner = value + if pre_p is not None and slot_key not in _slot_refs: + # The committed literal is the write this promotion serves: hand it to + # the structural-consumption arm so value-equality is judged on it. + ref = _meta_promote_slot( + slot_key, + pre_p, + target_name=str(key), + filename=filename, + lineno=lineno, + promoted_value=cur_p, + ) + if ref is None: + return False + _pyir_emit_store(_emit_constant_for_ref(ref, cur_p), ref) + staged = _load_as_dsl(ref, place=slot_key) + _pyir_holder_store(holder, key, staged) + else: + # The slot already carries staged state: the literal arrives through + # the instrumented write choke unchanged. + old = ( + None + if pre is _PYIR_BOUNDARY_MISSING + else pyir_read(label, pre, owner=holder, slot_name=key) + ) + result = pyir_assign( + label, old, cur, filename, lineno, owner=holder, slot_name=key + ) + if result is not cur: + _pyir_holder_store(holder, key, result) + # Recorded only after the commit succeeded: a bailed flip must leave no + # refusal residue on its guard places. + for hop_key, entry in recorded: + _PYIR_BOUNDARY_FLIP_GUARD_READS.setdefault(hop_key, entry) + log().info( + "[pyir boundary] committed literal store '%s' as a loop-carried cell " + "at the call site", + label, + ) + return True + + +def _pyir_boundary_replay( + holder: Any, + key: Any, + pre: Any, + cur: Any, + filename: str, + lineno: int, + callee: Any = None, +) -> None: + """Replay one observed write through the instrumented-write choke -- exactly + the ``holder.key = cur`` lifecycle, every tier and refusal unchanged.""" + label = _pyir_boundary_label(holder, key) + # META->META in staged CF: irreducible in a staged LOOP (refuse loudly), + # M2S-promotable in a loop-free ``if``; plain-Python shapes are exempt. + pre_p = ( + _pyir_unwrap_meta_primitive(pre) if pre is not _PYIR_BOUNDARY_MISSING else None + ) + cur_p = _pyir_unwrap_meta_primitive(cur) + # A raw-meta CURRENT value over a STAGED pre-value is the same invisible + # arithmetic: the advance did not come through the cell. + _pre_is_staged = False + if pre_p is None and pre is not _PYIR_BOUNDARY_MISSING: + try: + _pre_is_staged = _is_staged_value(pre) + except Exception: + _pre_is_staged = False + if ( + (pre_p is not None or _pre_is_staged) + and cur_p is not None + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + ): + slot_key = _make_slot_key(None, holder, key) + if slot_key is not None: + # In a staged LOOP both raw values meta = the advance bypassed the + # cell: refuse, unless a guard-stable LITERAL flip is proven. + if _innermost_enclosing_loop_op_at_ip() is not None: + if _pyir_boundary_loop_carry_flip( + callee, holder, key, pre, pre_p, cur, cur_p, label, filename, lineno + ): + return + raise DSLUserCodeError( + DiagId.BOUNDARY_META_LOOP_CARRY, + filename=filename, + lineno=lineno, + name=str(key), + ) + if pre_p is not None and slot_key not in _slot_refs: + # The observed post-call value is the write this promotion + # serves: value-equality at the structural arm judges it. + ref = _meta_promote_slot( + slot_key, + pre_p, + target_name=str(key), + filename=filename, + lineno=lineno, + promoted_value=cur_p, + ) + if ref is not None: + _pyir_emit_store(_emit_constant_for_ref(ref, cur_p), ref) + # F-SPEC write amendment: re-entry verifies the trace-exit + # payload of the promoted place, as on the choked path. + _pyir_spec_record_write(slot_key, cur_p) + staged = _load_as_dsl(ref, place=slot_key) + _pyir_holder_store(holder, key, staged) + log().info( + "[pyir boundary] promoted meta write '%s' at the call " + "site (loop-free staged if)", + label, + ) + return + if pre is _PYIR_BOUNDARY_MISSING: + old = None + else: + old = pyir_read(label, pre, owner=holder, slot_name=key) + value = cur + if type(value) is dict: + value = _pyir_adopt_dict_value(holder, key, value, label=label) + elif type(value) is list: + value = _pyir_adopt_list_value(holder, key, value, label=label) + result = pyir_assign( + label, old, value, filename, lineno, owner=holder, slot_name=key + ) + # Meta first-def from a plain callee: record the binding position (a + # staged first-def records through the assign choke). + if pre is _PYIR_BOUNDARY_MISSING and not _is_staged_value(value): + _pyir_record_cf_attr_first_def(holder, key) + # A FIRST-DEF from a plain callee in a staged LOOP has unprovable + # multiplicity: lift the keep-meta exemption so later changes fail-close. + if ( + pre is _PYIR_BOUNDARY_MISSING + and is_inside_staged_cf() + and _innermost_enclosing_loop_op_at_ip() is not None + ): + try: + _fd_slot = _make_slot_key(None, holder, key) + if _fd_slot is not None: + _slot_first_def_inside_cf[_fd_slot] = False + except Exception: + pass + if result is not cur: + _pyir_holder_store(holder, key, result) + # F-SPEC: the value left in storage (not the assign input) is what a + # re-entry re-derefs; re-amend the rows to the stored object. + _spec_place = _make_slot_key(None, holder, key) + if _spec_place is not None: + _pyir_spec_record_write(_spec_place, result) + log().info( + "[pyir boundary] committed un-instrumented write '%s' at the call site", + label, + ) + + +def _pyir_boundary_commit( + callee: Any, + snapshot: ( + "tuple[list[tuple[Any, Any, Any]], dict[int, tuple[Any, frozenset]], set[int]]" + ), + sc_rhs: bool = False, + result: Any = None, +) -> None: + """Diff the snapshot against the post-call state and replay each change, + deduped by leaf slot; deletions stay Python, *sc_rhs* effects refuse; a + *result* object born in the callee enters the identity domain (birth).""" + func = getattr(callee, "__func__", callee) + code = getattr(func, "__code__", None) + filename = getattr(code, "co_filename", "") + lineno = getattr(code, "co_firstlineno", 0) + leaves, key_sets, visited = snapshot + for holder, key, pre in leaves: + cur = _pyir_boundary_current_value(holder, key) + if cur is _PYIR_BOUNDARY_MISSING: + continue + if not _pyir_boundary_leaf_changed(holder, key, pre, cur): + continue + # The callee replaced this leaf: witness the pre-call binding so trace + # close can restore it if the leaf is left holding a staged wrapper. + _pyir_record_host_restore(holder, key, pre) + if sc_rhs: + mv = None + try: + mv = _get_slot_mv(holder, key) + except Exception: + mv = None + if is_inside_staged_cf() or (mv is not None and mv.ref is not None): + _pyir_boundary_sc_refuse(key, filename, lineno) + continue # plain trace scope: Python already has the value + _pyir_boundary_replay( + holder, + key, + pre, + cur, + filename, + lineno, + callee=callee, + ) + # Attributes / entries ADDED by the callee replay as first defs, entering + # the place model exactly like an instrumented first write. + for _oid, (holder, pre_keys) in key_sets.items(): + if isinstance(holder, _WatchedDict) or isinstance(holder, _WatchedList): + # Adopted containers intercept their own writes object-level. + continue + if isinstance(holder, dict): + post_keys = list(dict.keys(holder)) + elif isinstance(holder, list): + post_keys = list(range(list.__len__(holder))) + elif isinstance(holder, type): + post_keys = list(_pyir_class_facts.class_storage_items(holder).keys()) + else: + items_map = _instance_storage_items(holder) + if items_map is None: + continue + post_keys = [ + k + for k in items_map.keys() + if not (isinstance(k, str) and k.startswith("__")) + ] + for key in post_keys: + if key in pre_keys: + continue + cur = _pyir_boundary_current_value(holder, key) + if cur is _PYIR_BOUNDARY_MISSING: + continue + if sc_rhs: + if is_inside_staged_cf(): + _pyir_boundary_sc_refuse(key, filename, lineno) + continue + # A dict/list key CREATED by a plain callee inside dynamic staged + # CF is a key-set change -- the same declared fact the watched + # structural chokes refuse: a created entry has no pre-region + # cell, so it cannot follow the region's runtime predicate. + # Replaying it as a first def would bake it unconditionally. + if ( + isinstance(holder, (dict, list)) + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + ): + raise DSLUserCodeError( + ( + DiagId.CONTAINER_DICT_KEY_SET_MUTATED + if isinstance(holder, dict) + else DiagId.CONTAINER_LIST_SHAPE_MUTATED + ), + var=_pyir_boundary_label(holder, key), + detail=f"key {key!r} created by a plain (non-jit) callee", + ) + _pyir_boundary_replay( + holder, + key, + _PYIR_BOUNDARY_MISSING, + cur, + filename, + lineno, + callee=callee, + ) + # Birth-at-boundary-return: a returned object born in the callee enters + # the ledger's identity domain here (token mint + born-class stamp). + if result is not None and not sc_rhs: + _pyir_boundary_record_ctor_birth(result, visited) + + +def _pyir_boundary_record_ctor_birth(result: Any, visited: "set[int]") -> None: + """Mint the owner token (which stamps F-CLASS) for each object in *result* + NOT reachable at snapshot time and never tokenized -- the birth event. + Leaf publication stays with the binding chokes: the caller's bind roots the + object at its binding position and publishes its places there (LAW 1/3); + a boundary-side leaf replay would pre-freeze self-rooting and + double-publish content the trace already records.""" + + def _collect(obj: Any) -> None: + if obj is None or type(obj) in _PYIR_BOUNDARY_PASSTHROUGH_TYPES: + return + if isinstance(obj, (ir.Value, type, _types.ModuleType, _WatchedM)): + return + if id(obj) in visited: + return # reachable pre-call: not a birth + visited.add(id(obj)) + if isinstance(obj, (_WatchedDict, _WatchedList)): + return # adopted containers instrument their own writes + if isinstance(obj, (dict, list, tuple)): + vals = list(dict.values(obj)) if isinstance(obj, dict) else list(obj) + for v in vals: + _collect(v) + return + if _is_staged_value(obj): + return + if _pyir_lookup_owner_token(obj) is not None: + return # already a known owner: its writes are the diff's domain + items = _instance_storage_items(obj) + if items is None: + return + # S4 floor at birth: a user-defined `__del__` fires at a GC-determined + # instant, which has no binding position in the traced program. + _pyir_refuse_del_finalizer_owner(obj) + _owner_token(obj) # birth mint (stamps F-CLASS) + for v in list(items.values()): + _collect(v) + + _collect(result) + + +def _verify_no_used_poison(module: "ir.Module") -> None: + """Raise for a poison/placeholder-init ref read with no dominating store; the + analysis runs in C++ (``pyir.find_first_used_poison``), this shim only raises.""" + if pyir is None: + return + from . import pyir_core as _pyir_core + + if _pyir_core._POISON_EMITTED == 0: + # No poison / stamped placeholder init was produced since the last + # verify boundary, so the module cannot contain one: skip the scan. + return + # Consume the counter BEFORE the scan so the raising path does not leak it; + # a stale count from a failed trace only costs one spurious scan. + _pyir_core._POISON_EMITTED = 0 + finder = getattr(pyir, "find_first_used_poison", None) + if finder is None: + return + result = finder(module.operation) + if result is None: + return + filename, lineno, loc_str = result + if filename is None and loc_str: + import re + + m = re.search(r'"([^"]+\.py)"\s*:\s*(\d+)', loc_str) + if m: + filename, lineno = m.group(1), int(m.group(2)) + raise DSLUserCodeError( + DiagId.SCOPE_READ_NEVER_SET, filename=filename, lineno=lineno + ) + + +# Static export surface (regenerate with scripts/gen_pyir_all.py): the +# module's OWN names -- its namespace minus what lower chain modules +# already export and minus the generator's INTERNAL_ONLY registry. +__all__ = [ + "_bisect", + "_collections", + "_dataclasses", + "_dis", + "_heapq", + "_operator", + "_os", + "_types", + "_pyir_boundary_module_is_user", + "_pyir_boundary_consume_meta_arg", + "_pyir_delete_attr", + "_pyir_import_guard", + "_PYIR_GLOBAL_WRITE_SCANS", + "_pyir_callee_global_write", + "_pyir_global_write_guard", + "_pyir_global_write_watch", + "_pyir_global_write_check", + "_pyir_import_record", + "_pyir_routed_vars", + "_pyir_traced_locals", + "_pyir_routed_dict_view", + "_pyir_routed_clevel_mutator", + "_PYIR_CLEVEL_MUTATOR_ROUTES", + "_PYIR_DEQUE_BOUND_MUTATORS", + "_pyir_call_boundary_", + "_PYIR_SETATTR_OVERRIDES", + "_pyir_user_setattr_override", + "_pyir_is_property_store", + "_pyir_property_store", + "_pyir_inplace_binop", + "_pyir_boundary_value_tracked", + "_pyir_ctx_boundary", + "_pyir_boundary_bind_storage", + "_pyir_boundary_stage_taken_default_reads", + "_pyir_boundary_restore_meta_reads", + "_pyir_enclosing_while_cond_is_baked", + "_verify_no_used_poison", +] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_class_facts.py b/python/CuTeDSL/cutlass/base_dsl/pyir_class_facts.py index f3a2465e44..362738927f 100644 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_class_facts.py +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_class_facts.py @@ -1,4 +1,4 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: LicenseRef-NvidiaProprietary # # Use of this software is governed by the terms and conditions of the @@ -10,6 +10,1017 @@ # is strictly prohibited. -"""PyIR runtime -- class-facts layer; see facade for the public surface.""" +"""Decoration-time per-class ``self.`` write-fact registry: modules +parse ONCE at decoration/preprocess scope, never inside a trace-time walk.""" -__all__: list = [] +import ast +import hashlib +import inspect +import os +import sys +import tempfile +import textwrap +import types +from typing import Any, Optional + +from .pyir_state import _Sentinel + +# (module_name, def qualname, first lineno) -> {attr: (literal_values, +# guard_paths)} for defs whose ``attr`` writes are ALL guard-enumerable +# literal-constant stores. The lineno keeps qualname TWINS (a re-defined +# name shares the qualname) from colliding onto one row. +_LITERAL_WRITE_FACTS: "dict[tuple[str, str, int], dict[str, tuple[tuple, tuple]]]" = {} +# (module_name, def qualname, first lineno) -> (receiver attr writes, receiver +# method calls, root-first-arg free calls, other free calls, incomplete); +# exact-identity keyed rows for every def. ``incomplete`` is True when the +# write set is not provably complete: an opaque call, or a deep receiver-rooted +# write target (``self.sub.n``) the one-hop name schema cannot express. +_DEF_FACTS: "dict[tuple[str, str, int], tuple[frozenset, frozenset, frozenset, frozenset, bool]]" = {} +# Module names recorded by a decoration but not yet parsed. +_PENDING_MODULES: "list[str]" = [] +_PENDING_SET: "set[str]" = set() +# Modules already parsed (or unparseable -- both final). +_MATERIALIZED_MODULES: "set[str]" = set() +# Set on the first PyIR preprocessing session; from then on decorations +# materialize their module eagerly (decoration scope) instead of pending. +_PYIR_ACTIVE: bool = False +# DSL-declared module/package prefixes; re-resolved against ``sys.modules`` at +# every drain so a submodule imported later materializes at the next session. +_DECLARED_PREFIXES: "set[str]" = set() + + +def _receiver_attr_write_names( + fn: "ast.FunctionDef | ast.AsyncFunctionDef", +) -> "tuple[frozenset[str], bool]": + """``(attr_names, enumerable)``: attribute names *fn* assigns on its + first-parameter receiver, nested defs included (a nested impl's write + attributes to the enclosing callable). ``enumerable`` is False when a + write target is a deep receiver-rooted path (``self.sub.n``).""" + args = fn.args.posonlyargs + fn.args.args + if not args: + return frozenset(), True + recv = args[0].arg + found: "set[str]" = set() + enumerable = True + for node in ast.walk(fn): + targets: "list[ast.expr]" = [] + if isinstance(node, ast.Assign): + targets = node.targets + elif isinstance(node, (ast.AugAssign, ast.AnnAssign)): + targets = [node.target] + elif isinstance(node, (ast.For, ast.AsyncFor)): + targets = [node.target] + elif isinstance(node, (ast.With, ast.AsyncWith)): + targets = [ + item.optional_vars + for item in node.items + if item.optional_vars is not None + ] + for t in targets: + stack = [t] + while stack: + n = stack.pop() + if isinstance(n, (ast.Tuple, ast.List)): + stack.extend(n.elts) + elif isinstance(n, ast.Starred): + stack.append(n.value) + elif isinstance(n, ast.Attribute): + if isinstance(n.value, ast.Name) and n.value.id == recv: + found.add(n.attr) + continue + base = n.value + while isinstance(base, (ast.Attribute, ast.Subscript)): + base = base.value + if isinstance(base, ast.Name) and base.id == recv: + enumerable = False + return frozenset(found), enumerable + + +def _def_call_shapes( + fn: "ast.FunctionDef | ast.AsyncFunctionDef", +) -> "tuple[frozenset[str], frozenset[str], frozenset[str], bool]": + """Call shapes of *fn* for the per-def transitive-write closure: + ``(receiver_method_calls, root_first_arg_free_calls, other_free_calls, + has_opaque_calls)``. A free call joins the root set only when the + receiver is provably its first parameter (first positional argument is + the bare receiver name and the receiver escapes into no other position) — + the aliasing fact the closure needs to attribute the callee's + first-parameter writes to the root.""" + args = fn.args.posonlyargs + fn.args.args + recv = args[0].arg if args else None + methods: "set[str]" = set() + frees_root: "set[str]" = set() + frees_other: "set[str]" = set() + opaque = False + for node in ast.walk(fn): + if not isinstance(node, ast.Call): + continue + f = node.func + if isinstance(f, ast.Name): + first = node.args[0] if node.args else None + root_first = ( + recv is not None and isinstance(first, ast.Name) and first.id == recv + ) + root_elsewhere = recv is not None and any( + isinstance(a, ast.Name) and a.id == recv + for a in [ + *node.args[1:], + *(kw.value for kw in node.keywords), + ] + ) + if root_first and not root_elsewhere: + frees_root.add(f.id) + else: + frees_other.add(f.id) + elif ( + isinstance(f, ast.Attribute) + and isinstance(f.value, ast.Name) + and recv is not None + and f.value.id == recv + ): + methods.add(f.attr) + else: + opaque = True + return frozenset(methods), frozenset(frees_root), frozenset(frees_other), opaque + + +def _receiver_attr_chain(expr: "ast.expr", recv: str) -> "tuple[str, ...] | None": + """``self.a.b`` -> ``("a", "b")`` when rooted at the Name *recv*.""" + hops: "list[str]" = [] + node = expr + while isinstance(node, ast.Attribute): + hops.append(node.attr) + node = node.value + if isinstance(node, ast.Name) and node.id == recv and hops: + return tuple(reversed(hops)) + return None + + +def _guard_attr_paths(expr: "ast.expr", recv: str) -> "set[tuple[str, ...]] | None": + """Receiver-rooted attribute paths a guard expression reads, or ``None`` + when it reads anything the fact cannot enumerate.""" + if isinstance(expr, ast.Constant): + return set() + chain = _receiver_attr_chain(expr, recv) + if chain is not None: + return {chain} + if isinstance(expr, ast.BoolOp): + subs = [_guard_attr_paths(v, recv) for v in expr.values] + elif isinstance(expr, ast.UnaryOp): + subs = [_guard_attr_paths(expr.operand, recv)] + elif isinstance(expr, ast.Compare): + subs = [_guard_attr_paths(expr.left, recv)] + [ + _guard_attr_paths(c, recv) for c in expr.comparators + ] + elif isinstance(expr, ast.IfExp): + subs = [ + _guard_attr_paths(expr.test, recv), + _guard_attr_paths(expr.body, recv), + _guard_attr_paths(expr.orelse, recv), + ] + else: + return None + paths: "set[tuple[str, ...]]" = set() + for s in subs: + if s is None: + return None + paths.update(s) + return paths + + +def _collect_literal_attr_write_facts( + fn: "ast.FunctionDef | ast.AsyncFunctionDef", +) -> "dict[str, tuple[tuple, tuple]]": + """Per-def literal-write facts ``{attr: (literal_values, guard_paths)}``: + an attr qualifies only when its EVERY write is a literal store under + guards that read only enumerable receiver-rooted attribute chains.""" + pos_args = list(fn.args.posonlyargs) + list(fn.args.args) + recv = pos_args[0].arg if pos_args else None + if recv is None: + return {} + # A call can write receiver attributes only when the receiver escapes + # into it: a receiver-rooted method call, the bare receiver as an + # argument or alias, or a dynamic-access builtin. A call the receiver + # cannot reach (`Int32(self.x)`) leaves every write of this def visible + # to the scan below, so literal rows survive it; on any escape the def + # keeps no rows and the consumer keeps its loud refusal. (The def's own + # decorators/defaults evaluate at def time, not per call.) + consumed_recv: "set[int]" = set() + for st in fn.body: + for sub in ast.walk(st): + if isinstance(sub, ast.Call): + if ( + isinstance(sub.func, ast.Name) + and sub.func.id in _SURFACE_OPAQUE_CALLS + ): + return {} + if ( + isinstance(sub.func, ast.Attribute) + and _receiver_attr_chain(sub.func, recv) is not None + ): + # Receiver-rooted method call (self.m(), self.child.m()): + # the callee can reach the receiver, its writes are unseen. + return {} + if ( + isinstance(sub, ast.Attribute) + and isinstance(sub.value, ast.Name) + and sub.value.id == recv + ): + consumed_recv.add(id(sub.value)) + for st in fn.body: + for sub in ast.walk(st): + if ( + isinstance(sub, ast.Name) + and sub.id == recv + and id(sub) not in consumed_recv + ): + return {} # bare receiver escape (call arg, return, rebind) + + def _touched_attrs(t: "ast.expr", out: "set[str]") -> None: + if isinstance(t, (ast.Tuple, ast.List)): + for e in t.elts: + _touched_attrs(e, out) + return + if isinstance(t, ast.Starred): + _touched_attrs(t.value, out) + return + while isinstance(t, ast.Subscript): + t = t.value + if isinstance(t, ast.Attribute): + out.add(t.attr) + + def _stmt_touched_attrs(st: "ast.AST") -> "set[str]": + out: "set[str]" = set() + if isinstance(st, ast.Assign): + for t in st.targets: + _touched_attrs(t, out) + elif isinstance(st, (ast.AugAssign, ast.AnnAssign)): + _touched_attrs(st.target, out) + return out + + literals: "dict[str, list]" = {} + guard_map: "dict[str, set[tuple[str, ...]]]" = {} + max_write_lineno: "dict[str, int]" = {} + disqualified: "set[str]" = set() + min_exit_lineno: "list[int | None]" = [None] + + def _disqualify_writes_under(node: "ast.AST") -> None: + for sub in ast.walk(node): + disqualified.update(_stmt_touched_attrs(sub)) + + def _scan(stmts: "list[ast.stmt]", guards: "list[ast.expr]") -> None: + for st in stmts: + if isinstance(st, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + # A write hidden in a nested scope executes under call + # conditions the fact cannot enumerate. + _disqualify_writes_under(st) + continue + if isinstance(st, (ast.Return, ast.Raise, ast.Break, ast.Continue)): + if min_exit_lineno[0] is None or st.lineno < min_exit_lineno[0]: + min_exit_lineno[0] = st.lineno + continue + if isinstance(st, (ast.Assign, ast.AugAssign, ast.AnnAssign)): + touched = _stmt_touched_attrs(st) + if not touched: + continue + targets = st.targets if isinstance(st, ast.Assign) else [st.target] + t = targets[0] if len(targets) == 1 else None + value = getattr(st, "value", None) + if ( + isinstance(st, ast.AugAssign) # reads the (invisible) old value + or t is None + or not ( + isinstance(t, ast.Attribute) + and isinstance(t.value, ast.Name) + and t.value.id == recv + ) + or not isinstance(value, ast.Constant) + ): + disqualified.update(touched) + continue + attr = t.attr + paths: "set[tuple[str, ...]]" = set() + for g in guards: + gp = _guard_attr_paths(g, recv) + if gp is None: + disqualified.add(attr) + break + paths.update(gp) + else: + literals.setdefault(attr, []).append(value.value) + guard_map.setdefault(attr, set()).update(paths) + max_write_lineno[attr] = max( + max_write_lineno.get(attr, 0), st.lineno + ) + continue + if isinstance(st, ast.If): + _scan(st.body, guards + [st.test]) + _scan(st.orelse, guards + [st.test]) + continue + if isinstance(st, ast.While): + _scan(st.body + st.orelse, guards + [st.test]) + continue + if isinstance(st, (ast.For, ast.AsyncFor)): + _scan(st.body + st.orelse, guards + [st.iter]) + continue + # Any other compound statement (with/try/match) has execution + # conditions the fact cannot enumerate. + _disqualify_writes_under(st) + + _scan(fn.body, []) + result: "dict[str, tuple[tuple, tuple]]" = {} + for attr, lits in literals.items(): + if attr in disqualified: + continue + if min_exit_lineno[0] is not None and min_exit_lineno[0] < max_write_lineno.get( + attr, 0 + ): + continue # an early exit above a write acts as an unseen guard + result[attr] = (tuple(lits), tuple(sorted(guard_map.get(attr, ())))) + return result + + +def _walk_defs( + node: "ast.AST", + prefix: str, + out: "list[tuple[str, ast.FunctionDef | ast.AsyncFunctionDef]]", +) -> None: + """Collect every def under *node* with its runtime ``__qualname__`` + (class methods append ``.``; nested defs ``..``).""" + for child in ast.iter_child_nodes(node): + if isinstance(child, (ast.FunctionDef, ast.AsyncFunctionDef)): + qual = prefix + child.name + out.append((qual, child)) + _walk_defs(child, qual + "..", out) + elif isinstance(child, ast.ClassDef): + _walk_defs(child, prefix + child.name + ".", out) + else: + _walk_defs(child, prefix, out) + + +def _def_fact_key(func: Any) -> "tuple[str, str, int] | None": + """Exact def identity for the fact tables: (module, qualname, first line + of the decorated def block). A wrapper resolves through its declared + ``__wrapped__`` chain; a callable without a code object names no def, so + its lookups miss (consumers fail closed).""" + hops = 0 + while hasattr(func, "__wrapped__") and hops < 8: + func = func.__wrapped__ + hops += 1 + mod = getattr(func, "__module__", None) + qualname = getattr(func, "__qualname__", None) + code = getattr(func, "__code__", None) + if mod is None or qualname is None or code is None: + return None + return (mod, qualname, code.co_firstlineno) + + +def literal_attr_write_fact(func: Any, attr: str) -> "tuple[tuple, tuple] | None": + """The literal-write fact of ``func``'s exact def for ``attr``; ``None`` + without a qualifying row. Materializes ``func``'s module on first lookup + (memoized), so every fact read funnels through the single parse site.""" + key = _def_fact_key(func) + if key is None or not isinstance(attr, str): + return None + _materialize_module(key[0]) + row = _LITERAL_WRITE_FACTS.get(key) + if row is None: + return None + return row.get(attr) + + +# On-disk facts cache. WHAT is cached: one module's parse products, i.e. the +# rows the ast collectors above would install, keyed ``(qualname, lineno)`` +# per def: +# - literal rows (-> _LITERAL_WRITE_FACTS): for defs whose receiver-attr +# writes are all literal-constant stores, ``{attr: (literals, guards)}``. +# - def rows (-> _DEF_FACTS): every def's write/call shape: receiver attrs +# written, receiver methods called, free calls, and the incomplete flag. +# The trace-time attr-write fast paths read those registries; a cache hit +# installs the rows without re-parsing the module source (the expensive step). +# +# HOW it is stored: ``{dir}/{sha256(source)[:32]}.v{FORMAT}.facts`` holds the +# rows as a Python literal (``repr`` on write, ``ast.literal_eval`` on read, +# never pickle). Keyed by the exact source text the parser would see, so an +# edited module can never be served stale rows; unconfigured or disabled +# processes always parse. Bump the format when the collectors change shape. +_FACTS_CACHE_FORMAT = 1 +# None: not configured yet; "": configured off; else the cache directory. +_FACTS_CACHE_DIR: "list[str | None]" = [None] + + +def configure_facts_cache(cache_dir: str, enabled: bool) -> None: + """Configure the on-disk facts cache once per process (first DSL wins); + the rows are DSL-agnostic parse products, so one shared store is sound.""" + if _FACTS_CACHE_DIR[0] is None: + _FACTS_CACHE_DIR[0] = cache_dir if enabled else "" + + +def _facts_cache_path(source: str) -> "str | None": + d = _FACTS_CACHE_DIR[0] + if not d: + return None + digest = hashlib.sha256(source.encode("utf-8", "replace")).hexdigest()[:32] + return os.path.join(d, f"{digest}.v{_FACTS_CACHE_FORMAT}.facts") + + +def _facts_cache_load( + source: str, +) -> "tuple[dict, dict] | None": + """Cached ``(literal_rows, def_rows)`` for the exact *source* text, keyed + ``(qualname, lineno)``, or ``None``; a missing, corrupt or unreadable file + reports a miss (the caller parses), never raises.""" + path = _facts_cache_path(source) + if path is None: + return None + try: + with open(path, "r", encoding="utf-8") as f: + payload = ast.literal_eval(f.read()) + if payload["format"] != _FACTS_CACHE_FORMAT: + return None + def_rows = { + key: ( + frozenset(w), + frozenset(m), + frozenset(fr), + frozenset(fo), + bool(inc), + ) + for key, (w, m, fr, fo, inc) in payload["defs"].items() + } + return payload["literal"], def_rows + except Exception: + return None + + +def _facts_cache_dump(source: str, literal_rows: dict, def_rows: dict) -> None: + """Best-effort write-through. The payload must round-trip as a Python + literal (a module with an exotic constant is skipped, it just re-parses); + concurrent writers are safe, the text lands via temp file + atomic rename.""" + path = _facts_cache_path(source) + if path is None: + return + payload = { + "format": _FACTS_CACHE_FORMAT, + "literal": literal_rows, + "defs": { + key: ( + tuple(sorted(w)), + tuple(sorted(m)), + tuple(sorted(fr)), + tuple(sorted(fo)), + bool(inc), + ) + for key, (w, m, fr, fo, inc) in def_rows.items() + }, + } + text = repr(payload) + try: + if ast.literal_eval(text) != payload: + return + except Exception: + return + tmp = None + try: + os.makedirs(os.path.dirname(path), exist_ok=True) + fd, tmp = tempfile.mkstemp(dir=os.path.dirname(path), suffix=".tmp") + with os.fdopen(fd, "w", encoding="utf-8") as f: + f.write(text) + os.replace(tmp, path) + except Exception: + if tmp is not None: + try: + os.unlink(tmp) + except OSError: + pass + + +def _materialize_module(module_name: str) -> None: + """Parse *module_name*'s source once and register facts; unparseable + source marks it materialized with no facts (consumers do not engage). + The parse is skipped when the facts cache holds rows for the exact + source text.""" + if module_name in _MATERIALIZED_MODULES: + return + _MATERIALIZED_MODULES.add(module_name) + module = sys.modules.get(module_name) + if module is None: + return + try: + source = textwrap.dedent(inspect.getsource(module)) + except (OSError, TypeError, SyntaxError, ValueError): + return + cached = _facts_cache_load(source) + if cached is not None: + for key, facts in cached[0].items(): + _LITERAL_WRITE_FACTS[(module_name, *key)] = facts + for key, row in cached[1].items(): + _DEF_FACTS[(module_name, *key)] = row + return + try: + tree = ast.parse(source) + except (OSError, TypeError, SyntaxError, ValueError): + return + # Exact-identity literal-write facts for every def (the call-boundary + # loop-carry flip's positive evidence; absence keeps the loud refusal). + defs: "list[tuple[str, ast.FunctionDef | ast.AsyncFunctionDef]]" = [] + _walk_defs(tree, "", defs) + literal_rows: "dict[tuple[str, int], dict]" = {} + def_rows: "dict[tuple[str, int], tuple]" = {} + for qualname, fnode in defs: + # Runtime ``co_firstlineno`` starts at the FIRST DECORATOR line. + first_lineno = min([fnode.lineno] + [d.lineno for d in fnode.decorator_list]) + key = (module_name, qualname, first_lineno) + facts = _collect_literal_attr_write_facts(fnode) + if facts: + _LITERAL_WRITE_FACTS[key] = facts + literal_rows[(qualname, first_lineno)] = facts + # Per-code-object row for the boundary pre-stage's transitive + # write-closure (exact runtime identity, never bare-name-keyed). + _methods, _frees_root, _frees_other, _opaque = _def_call_shapes(fnode) + _writes, _enumerable = _receiver_attr_write_names(fnode) + row = ( + _writes, + _methods, + _frees_root, + _frees_other, + _opaque or not _enumerable, + ) + _DEF_FACTS[key] = row + def_rows[(qualname, first_lineno)] = row + _facts_cache_dump(source, literal_rows, def_rows) + + +def ensure_module_materialized(module_name: "str | None") -> None: + """Materialize *module_name* through the single memoized parse site; an + already-materialized module is an O(1) set lookup.""" + if not module_name: + return + _materialize_module(module_name) + + +def record_jit_decoration(func: Any) -> None: + """Record the decorated *func*'s module: pending while PyIR is inactive + (non-PyIR flows never pay a parse), materialized eagerly once active.""" + module_name = getattr(func, "__module__", None) + if not isinstance(module_name, str) or not module_name: + return + if module_name in _MATERIALIZED_MODULES: + return + if _PYIR_ACTIVE: + _materialize_module(module_name) + return + if module_name not in _PENDING_SET: + _PENDING_SET.add(module_name) + _PENDING_MODULES.append(module_name) + + +def register_class_fact_modules(*prefixes: str) -> None: + """Declare DSL module/package prefixes: the call boundary classifies them + as substrate; their write facts materialize lazily at each fact lookup.""" + _DECLARED_PREFIXES.update(prefixes) + + +def on_pyir_preprocess_session_start() -> None: + """Materialize the pending decoration-recorded modules and mark PyIR + active; all other modules parse lazily at their first fact lookup.""" + global _PYIR_ACTIVE + _PYIR_ACTIVE = True + if not _PENDING_MODULES: + return + pending = _PENDING_MODULES[:] + _PENDING_MODULES.clear() + _PENDING_SET.clear() + for module_name in pending: + _materialize_module(module_name) + + +# cls -> the MRO class defining ``__getattr__`` (None when absent); class +# protocol facts are definition-time facts, memoized like the parse rows. +_DUNDER_DEFINER_ABSENT: Any = _Sentinel("dunder definer absent") +_DUNDER_DEFINER_FACTS: "dict[tuple[type, str], Optional[type]]" = {} + + +def _mro_definer(cls: type, dunder: str) -> "Optional[type]": + """Class-fact: the MRO class defining *dunder* for instances of *cls*, + ``None`` when absent; memoized per ``(cls, dunder)``.""" + cached = _DUNDER_DEFINER_FACTS.get((cls, dunder), _DUNDER_DEFINER_ABSENT) + if cached is not _DUNDER_DEFINER_ABSENT: + return cached + definer = None + for klass in inspect.getmro(cls): + if dunder in klass.__dict__: + definer = klass + break + _DUNDER_DEFINER_FACTS[(cls, dunder)] = definer + return definer + + +def getattr_fabrication_definer(cls: type) -> "Optional[type]": + """Class-fact: the MRO class defining ``__getattr__`` for instances of + *cls*, ``None`` when absent. A failed attribute lookup on such instances + FABRICATES a value with no storage slot (LangRef 3.12 section 3.3.2); for + a metaclass *cls* the fabricated reads are on its classes.""" + return _mro_definer(cls, "__getattr__") + + +def del_finalizer_definer(cls: type) -> "Optional[type]": + """Class-fact: the MRO class defining ``__del__`` for instances of *cls*, + ``None`` when absent. Finalization (LangRef 3.12 section 3.3.1, + ``object.__del__``) runs at a garbage-collection-determined instant, + which has no binding position in a traced program.""" + return _mro_definer(cls, "__del__") + + +_EXTRACT_SURFACE_ABSENT: Any = _Sentinel("extract surface absent") +# type -> frozenset of touched field names, or None (surface underivable). +_EXTRACT_SURFACE_FACTS: "dict[type, Optional[frozenset]]" = {} + +# Calls whose result can reach the receiver's storage without a syntactic +# ``self.`` spelling: the touched-field set cannot be proven complete. +_SURFACE_OPAQUE_CALLS = frozenset( + ("super", "getattr", "setattr", "delattr", "vars", "eval", "exec", "locals") +) + + +def _surface_def_ast(fn: Any) -> "Optional[ast.FunctionDef]": + """The plain-def AST of *fn* (source parsed, never executed); ``None`` + when there is no source or the object is not a plain def.""" + if not isinstance(fn, types.FunctionType): + return None + try: + tree = ast.parse(textwrap.dedent(inspect.getsource(fn))) + except Exception: + return None + node = tree.body[0] if tree.body else None + if not isinstance(node, ast.FunctionDef): + return None + return node + + +def _surface_collect( + cls: type, fn: Any, touched: "set[str]", seen: "set[int]", depth: int +) -> bool: + """Union *fn*'s receiver-attribute touches (every branch: reads, writes, + method calls) into *touched*; recurse into receiver methods and touched + properties so backing storage names join. ``False`` when the receiver + escapes the syntactic ``self.`` discipline (bare receiver use, + dynamic attribute access, ``super()`` delegation) or a touched protocol + leg cannot be followed: the set cannot be proven complete, so no fact + may be produced.""" + if callable(fn): + try: + fn = inspect.unwrap(fn) # decorator wrappers: follow the real def + except Exception: + return False + if not isinstance(fn, types.FunctionType): + return False + # Cycle key: the resolved function object (wrapper code objects are + # SHARED across distinct wrapped defs, so code identity under-visits). + if id(fn) in seen: + return True # already unioned (cycle / repeated helper) + if depth <= 0: + return False + seen.add(id(fn)) + node = _surface_def_ast(fn) + if node is None or not node.args.args: + return False + recv = node.args.args[0].arg + # Receiver Name nodes consumed as an Attribute root (the one licensed use). + consumed: "set[int]" = set() + names: "set[str]" = set() + for sub in ast.walk(node): + if isinstance(sub, ast.Call): + if isinstance(sub.func, ast.Name) and sub.func.id in _SURFACE_OPAQUE_CALLS: + return False + if ( + isinstance(sub, ast.Attribute) + and isinstance(sub.value, ast.Name) + and sub.value.id == recv + ): + consumed.add(id(sub.value)) + if not sub.attr.startswith("__"): + names.add(sub.attr) + for sub in ast.walk(node): + if isinstance(sub, ast.Name) and sub.id == recv and id(sub) not in consumed: + return False # bare receiver escape (call arg, return, rebind) + touched.update(names) + # Follow receiver methods and properties: their interior reads name the + # backing storage the protocol depends on. + for name in sorted(names): + static = inspect.getattr_static(cls, name, _EXTRACT_SURFACE_ABSENT) + if static is _EXTRACT_SURFACE_ABSENT or isinstance(static, type): + continue # instance storage / class data: the touch is the fact + if isinstance(static, (types.MemberDescriptorType, types.GetSetDescriptorType)): + continue # __slots__ storage: the touch is the fact + if isinstance(static, property): + for leg in (static.fget, static.fset): + if leg is not None and not _surface_collect( + cls, leg, touched, seen, depth - 1 + ): + return False + continue + if callable(static): + if not _surface_collect(cls, static, touched, seen, depth - 1): + return False + continue + if hasattr(type(static), "__get__"): + return False # a descriptor leg the walk cannot follow + return True + + +def declared_extract_surface(cls: type) -> "Optional[frozenset]": + """Class-fact: the instance-field names *cls*'s own + ``__extract_mlir_values__`` touches (reads feeding or selecting its + leaves, plus extraction-internal bookkeeping writes), or ``None`` when + the surface is underivable. + + Derived from the protocol itself, statically: the extraction def's AST + is unioned over every branch (a superset of any one execution), receiver + method calls recurse, and a touched property contributes its getter's + and setter's backing reads. The extraction body is NEVER executed -- + observation must not run user code. A shape the walk cannot prove + complete (bare receiver escape, dynamic attribute access, ``super()``) + yields ``None``: no fact, so the consumer admits the write.""" + cached = _EXTRACT_SURFACE_FACTS.get(cls, _EXTRACT_SURFACE_ABSENT) + if cached is not _EXTRACT_SURFACE_ABSENT: + return cached + fn = None + for klass in inspect.getmro(cls): + fn = klass.__dict__.get("__extract_mlir_values__") + if fn is not None: + break + touched: "set[str]" = set() + result: "Optional[frozenset]" = None + try: + if _surface_collect(cls, fn, touched, set(), 8): + result = frozenset(touched) + except Exception: + result = None + _EXTRACT_SURFACE_FACTS[cls] = result + return result + + +_HASH_OVERRIDE_ABSENT: Any = _Sentinel("hash override absent") +_HASH_OVERRIDE_FACTS: "dict[type, Optional[type]]" = {} + + +def staged_wrapper_hash_override(cls: type) -> "Optional[type]": + """Class-fact: the MRO class whose ``__hash__`` SHADOWS a staged + wrapper's own ``__hash__`` for instances of *cls*, ``None`` when absent. + A wrapper's ``__hash__`` is a consumption/identity witness; a subclass + override silently answers in its place, so hashing bypasses the witness. + ``__hash__ = None`` is exempt: an unhashable subclass fails loudly.""" + cached = _HASH_OVERRIDE_FACTS.get(cls, _HASH_OVERRIDE_ABSENT) + if cached is not _HASH_OVERRIDE_ABSENT: + return cached + result: "Optional[type]" = None + mro = inspect.getmro(cls) + definer_idx = -1 + for i, klass in enumerate(mro): + if "__hash__" in klass.__dict__: + definer_idx = i + break + if ( + definer_idx >= 0 + and mro[definer_idx] is not object + and callable(mro[definer_idx].__dict__.get("__hash__")) + ): + for klass in mro[definer_idx + 1 :]: + if "__hash__" not in klass.__dict__: + continue + # The entry the override shadows: a violation only when it is a + # staged wrapper's own (callable) hash. + if ( + klass is not object + and callable(klass.__dict__.get("__hash__")) + and callable(getattr(klass, "ir_value", None)) + ): + result = mro[definer_idx] + break + _HASH_OVERRIDE_FACTS[cls] = result + return result + + +def class_storage_items(cls: type) -> "dict[str, Any]": + """Class-level DATA attributes of *cls* (own ``__dict__`` only): the + class-attr places the boundary diff observes, with the CLASS object as + owner (F-SHAPE). Descriptors (functions, properties, classmethods, slot + members) and nested classes are protocol, not storage.""" + items: "dict[str, Any]" = {} + for name, v in cls.__dict__.items(): + if not isinstance(name, str) or name.startswith("__"): + continue + if isinstance(v, type) or hasattr(type(v), "__get__"): + continue + items[name] = v + return items + + +# Builtin predicates a transparent guard may call: pure, replay-stable. +_PURE_GUARD_CALLS = frozenset(("type", "isinstance", "len")) + + +def setattr_storage_transparent(fn: Any) -> bool: + """Can the ``__setattr__`` override *fn* be proven to store only its + UNMODIFIED value parameter (a storage redirect, e.g. into a dict field)? + Only a transparent override may stay on the rewritten-choke store path: + the choke's instrumentation re-binds REPLAY the override once per + re-bind, so every replayed expression -- guards, store targets, slot + keys -- must be side-effect-free, and the stored value must be the bare + parameter. ``False`` whenever the proof fails (no source, helper calls, + computed stores, effectful guards).""" + node = _surface_def_ast(fn) + if node is None or len(node.args.args) != 3: + return False + value_param = node.args.args[2].arg + + def _bare_value(expr: "ast.expr") -> bool: + return isinstance(expr, ast.Name) and expr.id == value_param + + def _simple(expr: "ast.expr") -> bool: + if isinstance(expr, (ast.Name, ast.Constant)): + return True + if isinstance(expr, ast.Attribute): + return _simple(expr.value) + if isinstance(expr, ast.Subscript): + return _simple(expr.value) and _simple(expr.slice) + if isinstance(expr, ast.Tuple): + return all(_simple(e) for e in expr.elts) + if isinstance(expr, ast.UnaryOp): + return _simple(expr.operand) + if isinstance(expr, ast.BoolOp): + return all(_simple(v) for v in expr.values) + if isinstance(expr, ast.Compare): + return _simple(expr.left) and all(_simple(c) for c in expr.comparators) + if isinstance(expr, ast.Call): + return ( + isinstance(expr.func, ast.Name) + and expr.func.id in _PURE_GUARD_CALLS + and not expr.keywords + and all(_simple(a) for a in expr.args) + ) + return False + + def _stmt_ok(st: "ast.stmt") -> bool: + if isinstance(st, ast.If): + return _simple(st.test) and all(_stmt_ok(s) for s in st.body + st.orelse) + if isinstance(st, (ast.Raise, ast.Pass)): + return True + if isinstance(st, ast.Return): + return st.value is None + if isinstance(st, ast.Assign): + return _bare_value(st.value) and all(_simple(t) for t in st.targets) + if isinstance(st, ast.Expr): + if isinstance(st.value, ast.Constant): + return True # docstring + call = st.value + return ( + isinstance(call, ast.Call) + and isinstance(call.func, ast.Attribute) + and call.func.attr == "__setattr__" + and isinstance(call.func.value, ast.Name) + and call.func.value.id == "object" + and len(call.args) == 3 + and not call.keywords + and _simple(call.args[0]) + and _simple(call.args[1]) + and _bare_value(call.args[2]) + ) + return False + + return all(_stmt_ok(s) for s in node.body) + + +# (def identity key, receiver type or None) -> (attr_writes, complete); +# memoized like the definer facts above (the substrate tables never re-parse). +_TRANSITIVE_WRITE_CACHE: "dict[tuple[tuple[str, str, int], Optional[type]], tuple[frozenset[str], bool]]" = {} + + +def transitive_write_facts(callee: Any) -> "tuple[frozenset[str], bool]": + """Receiver-attribute write set of *callee*'s exact def closed over its + resolvable call edges: ``(attr_writes, complete)``; under-approximates. + Memoized per (def identity, receiver type): the walk reads the receiver + only through ``type(recv).__mro__``, and the parse tables freeze at first + materialization, so the closure is a definition-time fact of that pair.""" + func = getattr(callee, "__func__", callee) + receiver = getattr(callee, "__self__", None) + root_key = _def_fact_key(func) + if root_key is None: + return frozenset(), False + cache_key = (root_key, type(receiver) if receiver is not None else None) + cached = _TRANSITIVE_WRITE_CACHE.get(cache_key) + if cached is not None: + return cached + ensure_module_materialized(root_key[0]) + root = _DEF_FACTS.get(root_key) + if root is None: + _TRANSITIVE_WRITE_CACHE[cache_key] = (frozenset(), False) + return frozenset(), False + writes: "set[str]" = set() + complete = True + seen: "set[tuple[str, str, int]]" = set() + stack: "list[tuple[tuple[str, str, int], Any, Any]]" = [(root_key, func, receiver)] + while stack: + key, fobj, recv = stack.pop() + if key in seen: + continue + seen.add(key) + row = _DEF_FACTS.get(key) + if row is None: + complete = False + continue + row_writes, row_methods, row_frees_root, row_frees_other, row_incomplete = row + writes.update(row_writes) + if row_incomplete: + complete = False + for m in row_methods: + target = None + if recv is not None: + for klass in type(recv).__mro__: + cand = klass.__dict__.get(m) + if cand is not None: + target = getattr(cand, "__func__", cand) + break + if not callable(target): + complete = False + continue + t_key = _def_fact_key(target) + if t_key is None: + complete = False + continue + ensure_module_materialized(t_key[0]) + stack.append((t_key, target, recv)) + fglobals = getattr(fobj, "__globals__", None) or {} + for fname in row_frees_root | row_frees_other: + target = fglobals.get(fname) + target = getattr(target, "__func__", target) + if not callable(target) or not hasattr(target, "__code__"): + # No def row and no syntactic attribute writes (builtin / C + # function); a dynamic setattr is the boundary diff's to observe. + continue + t_key = _def_fact_key(target) + if t_key is None: + continue + if fname in row_frees_other: + # The receiver is not provably the callee's first parameter + # at (at least) one call site: those writes belong to ANOTHER + # object, so merging them onto the root would promote + # never-written fields. The unattributable edge marks the + # set incomplete; a root-first call site of the same name + # still contributes below. + complete = False + if fname not in row_frees_root: + continue + ensure_module_materialized(t_key[0]) + if t_key not in _DEF_FACTS: + complete = False + continue + # A root-first-arg free call binds the ROOT as its first + # parameter; its row contributes writes and keeps the root + # receiver for its own method edges. + stack.append((t_key, target, recv)) + result = (frozenset(writes), complete) + _TRANSITIVE_WRITE_CACHE[cache_key] = result + return result + + +# Static export surface (regenerate with scripts/gen_pyir_all.py): the +# module's OWN names -- its namespace minus what lower chain modules +# already export and minus the generator's INTERNAL_ONLY registry. +__all__ = [ + "ast", + "hashlib", + "inspect", + "os", + "sys", + "tempfile", + "textwrap", + "types", + "Any", + "Optional", + "_Sentinel", + "_DECLARED_PREFIXES", + "literal_attr_write_fact", + "configure_facts_cache", + "_materialize_module", + "ensure_module_materialized", + "record_jit_decoration", + "register_class_fact_modules", + "on_pyir_preprocess_session_start", + "getattr_fabrication_definer", + "del_finalizer_definer", + "_EXTRACT_SURFACE_ABSENT", + "_EXTRACT_SURFACE_FACTS", + "_SURFACE_OPAQUE_CALLS", + "_surface_def_ast", + "_surface_collect", + "declared_extract_surface", + "_HASH_OVERRIDE_ABSENT", + "_HASH_OVERRIDE_FACTS", + "staged_wrapper_hash_override", + "class_storage_items", + "_PURE_GUARD_CALLS", + "setattr_storage_transparent", + "transitive_write_facts", +] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_cleanup.py b/python/CuTeDSL/cutlass/base_dsl/pyir_cleanup.py deleted file mode 100644 index fcde501b65..0000000000 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_cleanup.py +++ /dev/null @@ -1,218 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: LicenseRef-NvidiaProprietary -# -# Use of this software is governed by the terms and conditions of the -# NVIDIA End User License Agreement (EULA), available at: -# https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/license.html -# -# Any use, reproduction, disclosure, or distribution of this software -# and related documentation outside the scope permitted by the EULA -# is strictly prohibited. - - -"""PyIR runtime -- cleanup layer; see facade for the public surface.""" - -from .pyir_call_boundary import * # noqa: F401,F403 (re-export lower layers up the chain) - -from .._mlir import ir - -# -- BEGIN explicit imports for the type checker (do not edit the list by hand; -# it mirrors names the chain re-exports at runtime via the wildcard + dynamic -# ``__all__`` above, which a static type checker cannot evaluate -- so every -# name is also imported explicitly from the layer that DEFINES it). Purely -# additive: the wildcard import stays the runtime source of truth. -from .pyir_state import ( # noqa: F401 - Any, - DSLUserCodeError, - DiagId, - ub, -) -# -- END explicit imports for the type checker - - -def _verify_no_used_poison(module: "ir.Module") -> None: - """End-of-trace check: any ``ub.poison`` value with a real use is a bug. - - A poison-init ref is fine when no read ever consumes it — it's just a - placeholder waiting for a store. But if the IR contains a - ``pyir.load`` against a poison-init ref whose result feeds any - downstream op (printf, store, control-flow condition, etc.), the - load will return undefined data at runtime. We flag that case at - trace completion, before the C++ pass pipeline runs. - - Raises ``DSLUserCodeError`` with a Python-source location lifted from - the offending load's MLIR ``loc()`` attribute (set by the AST - preprocessor at the original Python access site). - """ - if ub is None: - return # ub dialect not built in -- no poison ops can exist - from . import pyir_core as _pyir_core - - if _pyir_core._POISON_EMITTED == 0: - # ``_make_poison_like`` is the only producer of ``ub.poison``; if it - # never ran since the last verify boundary the module cannot contain - # one. Skips the whole-module walk on every classic-mode (and most - # pyir) compiles. - return - # Consume the counter at the verify boundary: traces are sequential - # in-process, so every poison op built since the previous verify belongs - # to THIS module and is accounted for by the walk below. Consumed BEFORE - # the walk so the raising path (poison-read diagnostic) does not leak the - # count either. A failing trace that never reaches build_module (and so - # skips this verify) can leave a stale count behind -- that is fail-safe: - # the next verify walks once spuriously (sound, just slower) and re-arms - # the skip. - _pyir_core._POISON_EMITTED = 0 - - def _has_any_use(value: Any) -> bool: - try: - for _ in value.uses: - return True - except (AttributeError, TypeError): - return False - return False - - def _dominates(store_op: Any, load_op: Any) -> bool: - """True if ``store_op`` dominates ``load_op`` in the structured-region - sense: either same block with ``store_op`` earlier, or some ancestor - of ``load_op`` lives in ``store_op``'s block after ``store_op``. - - Walks up from ``load_op`` through ``op.block.region.owner`` (the - proper MLIR Python chain — ``Block.owner`` is not a region). When - an ancestor lands in ``store_op``'s block, checks whether - ``store_op`` precedes that ancestor (which is what dominance - requires for structured CF).""" - # ``OpView`` wrappers (StoreOp, LoadOp, ...) don't expose ``.block`` - # or ``.is_before_in_block`` directly — those live on the underlying - # ``Operation``. Normalise both inputs first. - store_raw = getattr(store_op, "operation", store_op) - load_raw = getattr(load_op, "operation", load_op) - - def _safe_block(op: Any) -> Any: - """Return ``op.block`` or None — guards the C++ assertion - ``Attached operation has null parent`` for module roots / - detached ops.""" - try: - if getattr(op, "parent", None) is None: - return None - return op.block - except Exception: - return None - - store_block = _safe_block(store_raw) - if store_block is None: - return False - cur = load_raw - depth = 0 - while cur is not None and depth < 64: - cur_block = _safe_block(cur) - if cur_block is None: - return False - try: - same_block = cur_block == store_block - except Exception: - same_block = False - if same_block: - if cur is store_raw: - return False - try: - return store_raw.is_before_in_block(cur) - except (AttributeError, TypeError): - return False - region = getattr(cur_block, "region", None) - if region is None: - return False - parent_op = getattr(region, "owner", None) - if parent_op is None or parent_op is cur: - return False - parent_op = getattr(parent_op, "operation", parent_op) - cur = parent_op - depth += 1 - return False - - def _load_has_dominating_store(ref_op: Any, load_op: Any) -> bool: - for use in ref_op.result.uses: - store = use.owner - if store is None or store.name != "pyir.store": - continue - if _dominates(store, load_op): - return True - return False - - offender = None # the offending pyir.load or other op consuming poison - offender_poison = None # the originating ub.poison ir.Value, if known - - def _walk(op: Any) -> None: - # Same walk, but also record the offending poison value when found. - nonlocal offender, offender_poison - if offender is not None: - return - if op.name == "ub.poison": - for use in op.result.uses: - user = use.owner - if user is None: - continue - if user.name == "pyir.ref": - for ref_use in user.result.uses: - ref_user = ref_use.owner - if ( - ref_user is not None - and ref_user.name == "pyir.load" - and _has_any_use(ref_user.result) - and not _load_has_dominating_store(user, ref_user) - ): - offender = ref_user - offender_poison = op.result - return - elif user.name == "pyir.load": - if _has_any_use(user.result): - offender = user - offender_poison = op.result - return - else: - for r in user.results: - if _has_any_use(r): - offender = user - offender_poison = op.result - return - for region in op.regions: - for block in region.blocks: - for sub in block.operations: - _walk(sub) - if offender is not None: - return - - _walk(module.operation) - if offender is None: - return - - # First try to recover Python file/line from the recorded poison - # source (set by _make_poison_like at creation). Fall back to parsing - # the offending load's MLIR loc() if that map was wiped. - filename, lineno = None, None - # Prefer the attributes stamped onto the originating ub.poison op by - # ``_record_poison_source``; fall back to the offender's MLIR loc(). - if offender_poison is not None: - try: - poison_op = offender_poison.owner - attrs = poison_op.attributes - file_attr = attrs["pyir.poison_src_file"] - line_attr = attrs["pyir.poison_src_line"] - filename = str(file_attr).strip('"') - lineno = int(str(line_attr).split()[0]) - except (KeyError, AttributeError, ValueError, TypeError): - pass - if filename is None: - import re - - loc_str = str(offender.location) if hasattr(offender, "location") else "" - m = re.search(r'"([^"]+\.py)"\s*:\s*(\d+)', loc_str) - if m: - filename, lineno = m.group(1), int(m.group(2)) - - raise DSLUserCodeError( - DiagId.SCOPE_READ_NEVER_SET, filename=filename, lineno=lineno - ) - - -__all__ = [name for name in list(globals()) if not name.startswith("__")] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_core.py b/python/CuTeDSL/cutlass/base_dsl/pyir_core.py index fa4af14c6e..7c5da819a4 100644 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_core.py +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_core.py @@ -10,1316 +10,1164 @@ # is strictly prohibited. -"""PyIR runtime -- core layer; see facade for the public surface.""" +"""PyIR runtime -- core layer.""" + +import enum as _enum +import functools as _functools +import gc +import sys as _sys +import types +import inspect +import weakref + +from typing import NoReturn from .pyir_state import * # noqa: F401,F403 (re-export lower layers up the chain) +from .pyir_state import _Sentinel + + +# Staleness tags stamped on every ``MutableValue.load`` result: the slot's load +# counter and the staged-CF depth the load was emitted at. +_PYIR_LOAD_VERSION_ATTR = "_pyir_load_version" +_PYIR_LOAD_DEPTH_ATTR = "_pyir_load_depth" +_PYIR_REGION_ENTRY_ATTR = "_pyir_region_entry" +# The ref's write-epoch at load time; ``_pyir_emit_store`` bumps it, so it also +# counts stores that bypassed the tagging ``MutableValue``. +_PYIR_REF_EPOCH_ATTR = "_pyir_ref_epoch" + +# Per-ref write-epoch registry, keyed by the raw ``pyir.ref`` value; trace-scoped. +_REF_WRITE_EPOCH: "dict[Any, int]" = {} + +# Closure-cell meta reads observed at a call boundary: slot key -> (name, +# callee file, callee line); promotion consults it to refuse loudly. Trace-scoped. +_PYIR_BOUNDARY_META_CELL_READS: "dict[Any, tuple[str, str, int]]" = {} + +# Structural (``__index__``) consumptions of a meta place inside staged CF: slot +# key -> (label, file, line); promotion consults it to refuse loudly. Trace-scoped. +_PYIR_STRUCTURAL_META_CONSUMPTIONS: "dict[Any, tuple[str, str, int]]" = {} + +# Where each structural consumption was witnessed: slot key -> (block, enclosing +# loop op). Kept beside ``_PYIR_STRUCTURAL_META_CONSUMPTIONS`` rather than widening +# its tuple, which several unpackers destructure positionally. Used to tell a +# re-seeded place (the loop body re-runs its binding before the fold is reached +# again, so the bake stays valid) from one whose value genuinely carries. +# Trace-scoped. +_PYIR_STRUCTURAL_META_CONSUMPTION_SITES: "dict[Any, tuple[Any, Any]]" = {} + +# Meta writes admitted inside a staged-``if`` arm of the slot's own birth +# region (a per-iteration reset counter): slot key -> (arm block, file, line). +# In-arm consumptions fold their program point's value exactly; a consumption +# OUTSIDE the arm is path-dependent (the arm may not run) and refuses loudly. +# The next same-depth reset retires the mark. Trace-scoped. +_PYIR_ARM_LOCAL_META_WRITES: "dict[Any, tuple[Any, str, int]]" = {} + +# Guard places a committed boundary loop-carry flip relied on: slot key -> +# (guard attr, flipped slot, callee file, callee line); promotion consults it. +_PYIR_BOUNDARY_FLIP_GUARD_READS: "dict[Any, tuple[str, str, str, int]]" = {} + +# Host-restore ledger: (id(holder), slot) -> (holder, slot, pre-staging scalar +# binding, minting-context id). Trace close re-points every recorded place still +# bound to THIS compilation's wrapper back to that binding; a leftover without a +# record keeps the stale-epoch refusal. Trace-scoped. +_PYIR_HOST_RESTORE: "dict[tuple[int, Any], tuple[Any, Any, Any, int]]" = {} + +# Literal-backed STAGED predicate folds inside an enclosing loop: id(src wrapper) +# -> (payload repr, file, line); staged write chokes refuse on a hit. Trace-scoped. +_PYIR_STAGED_LITERAL_FOLD_WITNESSES: "dict[int, tuple[str, str, int]]" = {} +_PYIR_STAGED_LITERAL_FOLD_KEEPALIVE: "list[Any]" = [] +_PYIR_TRACE_KEEPALIVES.append(_PYIR_STAGED_LITERAL_FOLD_KEEPALIVE) +# Promotion rewrite map: baked ``arith.constant`` -> the position-correct +# ``pyir.load`` that replaced its uses; new consumers follow it. Trace-scoped. +_META_CONST_REPLACEMENTS: "dict[Any, Any]" = {} + + +def _ref_write_epoch(ref: "Any") -> int: + """Current write-epoch of *ref* (0 = never stored through the choke).""" + if ref is None: + return 0 + try: + return _REF_WRITE_EPOCH.get(ref, 0) + except TypeError: + return 0 # unhashable ref stand-in -- no epoch tracking -# -- BEGIN explicit imports for the type checker (do not edit the list by hand; -# it mirrors names the chain re-exports at runtime via the wildcard + dynamic -# ``__all__`` above, which a static type checker cannot evaluate -- so every -# name is also imported explicitly from the layer that DEFINES it). Purely -# additive: the wildcard import stays the runtime source of truth. -from .pyir_state import ( # noqa: F401 - Any, - DSLRuntimeError, - Optional, - _MAX_FOLD_WITNESSES_PER_SLOT, - _MODULE_OPS, - _NON_DOT_FUNC_ENTRY_OPS, - _PYIR_SLOT_FALLBACK, - _SCF_REGION_NAMES, - _SLOT_STORE_ATTR, - _fold_witnesses, - _is_staged_value, - _meta_uses, - _pinned_owners, - _pyir_fn_id_stack, - _slot_first_def_block, - _slot_first_def_depth, - _slot_first_def_depth_any, - _slot_first_def_inside_cf, - _slot_mvs, - _slot_pending_store, - _slot_refs, - inspect, - ir, - log, - pyir, - ub, -) -# -- END explicit imports for the type checker -import operator -import sys -from typing import Callable # noqa: F401 (used in _BINOP_OPS annotation) -from types import FrameType # noqa: F401 (used in _fold_consumer_is_user_code) +def _get_load_version(v: "Any") -> "Optional[int]": + """Return *v*'s load-version tag, or ``None`` if it is not a load.""" + return getattr(v, _PYIR_LOAD_VERSION_ATTR, None) -def _make_slot_key( - target_name: "str | None", - owner: Any = None, - slot_name: Any = None, - fn_key: Any = None, -) -> Any: - """Build the key under which a D1 slot is tracked. +def _get_load_depth(v: "Any") -> "Optional[int]": + """Return *v*'s staged-CF load-depth tag, or ``None`` if untagged.""" + return getattr(v, _PYIR_LOAD_DEPTH_ATTR, None) - Locals: ``("name", target_name)`` - Attributes: ``("attr", id(owner), slot_name)`` - Returns ``None`` when neither a target name nor an owner/slot tuple - is available -- retroactive promotion is disabled for that call. - """ - if owner is not None and slot_name is not None: - # Keep the owner alive so its address can't be reused while we key on it. - # This is total: the key is id(owner) (an int), so the owner is never - # hashed and only stored by reference -- it can't fail for any owner, - # including unhashable ones like dict/list. - _pinned_owners[id(owner)] = owner - return ("attr", id(owner), slot_name) - if target_name is not None: - # Frame-local key: qualify by the owning user function (fn_key) so two - # functions with a same-named local get DISTINCT slots instead of - # colliding on one shared ``pyir.ref``. fn_key is None outside a tracked - # function -> fall back to the name-only key. - if fn_key is not None: - return ("name", target_name, fn_key) - return ("name", target_name) - return None +def _get_load_epoch(v: "Any") -> "Optional[int]": + """Return *v*'s ref-write-epoch tag, or ``None`` if untagged.""" + return getattr(v, _PYIR_REF_EPOCH_ATTR, None) -def _cached_ir_value_dominates_ip( - value: "ir.Value", - ip: "Optional[ir.InsertionPoint]" = None, -) -> bool: - """Whether *value* is reachable from the target insertion point. - - The target is *ip* when given, otherwise ``InsertionPoint.current``. - Used to invalidate per-instance ``_WatchedM._cached_ir`` when the - cached ``arith.constant`` was materialised in a sibling staged-CF - region (and therefore would be a region escape if reused). When - the ``pyir`` python binding is unavailable (PyIR-disabled builds) - or there is no active insertion point, returns True so the legacy - behaviour is preserved. - """ - if pyir is None: - return True - try: - current_block = ( - ip if ip is not None else ir.InsertionPoint.current - ).block - except (RuntimeError, ValueError): - return True - return pyir.is_value_in_ancestor_region(value, current_block) +def _get_load_region_entry(v: "Any") -> "Optional[int]": + """Return *v*'s region-entry-clock tag, or ``None`` if untagged.""" + return getattr(v, _PYIR_REGION_ENTRY_ATTR, None) -def _emit_constant_at_ip( - value: Any, - *, - loc: "Optional[ir.Location]" = None, - ip: "Optional[ir.InsertionPoint]" = None, -) -> "ir.Value": - """Emit a FRESH ``arith.constant`` for a Python primitive at the current IP. - - Bypasses the ``@lru_cache_ir`` memoisation on - :func:`cutlass._mlir_helpers.arith.const` so each call returns a - distinct ``ir.Value``. D1 retroactive promotion relies on - per-slot ``replaceAllUsesWith``: if multiple slots share the same - cached SSA, promoting any one of them would clobber the others' - uses too (the printf-with-three-True-args mixing.py bug). - """ - from .._mlir.dialects import arith as _arith - from .._mlir import ir as _ir +def _tag_load(v: "Any", version: int, depth: int, epoch: int = 0) -> None: + """Tag *v* with the load *version*, staged-CF *depth*, ref write-*epoch*, + and the region-entry clock it was loaded at. - if isinstance(value, bool): - mlir_ty = _ir.IntegerType.get_signless(1) - return _arith.constant(mlir_ty, value, loc=loc, ip=ip) - if isinstance(value, int): - mlir_ty = _ir.IntegerType.get_signless(32) - return _arith.constant(mlir_ty, value, loc=loc, ip=ip) - if isinstance(value, float): - mlir_ty = _ir.F32Type.get() - return _arith.constant(mlir_ty, value, loc=loc, ip=ip) - # Fallback (rare): defer to the cached helper. - from .._mlir_helpers.arith import const as _arith_const + Best-effort: value types that reject attribute assignment simply go + untagged (no dedup), which is safe. + """ + object.__setattr__(v, _PYIR_LOAD_VERSION_ATTR, version) + object.__setattr__(v, _PYIR_LOAD_DEPTH_ATTR, depth) + object.__setattr__(v, _PYIR_REF_EPOCH_ATTR, epoch) + object.__setattr__(v, _PYIR_REGION_ENTRY_ATTR, _pyir_region_entry_clock()) + _pyir_note_holder_write(v) # one stamp covers the tag-attr cluster - return _arith_const(value, loc=loc, ip=ip) +def _pyir_adopt_stored_representative(mv: Any, wrapper: Any) -> None: + """Record *wrapper* (a rewrap of the just-stored raw) as the cell's canonical + representative and stamp the store tags (epoch + depth, no load-version).""" + try: + mv._value = wrapper + object.__setattr__(wrapper, _PYIR_REF_EPOCH_ATTR, _ref_write_epoch(mv.ref)) + object.__setattr__(wrapper, _PYIR_LOAD_DEPTH_ATTR, current_staged_cf_depth()) + object.__setattr__(wrapper, _PYIR_REGION_ENTRY_ATTR, _pyir_region_entry_clock()) + _pyir_note_holder_write(wrapper) # one stamp covers the tag-attr cluster + except (AttributeError, TypeError): + pass # wrapper rejects attrs -- no staleness tracking. -def _unwrap(value: Any) -> Any: - """Strip a ``_WatchedM`` wrapper (one level), or return *value* as-is.""" - if isinstance(value, _WatchedM): - return value.python_value - return value +def _is_func_boundary_op(op_name: str) -> bool: + """Return True if *op_name* is a function-like op that owns an SSA body + region where a ``pyir.ref`` may be hosted. -def _fold_consumer_is_user_code() -> bool: - """Whether the frame that triggered this fold is author code, not DSL glue. - - Gate for fold-witness recording. The predicate positions the preprocessor - can see -- ternaries, ``while`` conditions -- record their witness directly - through ``force_consumer`` and never reach here. What remains is IMPLICIT - consumption (``__bool__``/``__eq__``/``__index__``/``__int__``/``__format__``) - that CPython invokes from an arbitrary call depth with no source site to - instrument: e.g. the comparisons inside ``max``/``min``/``sum`` run in their - C implementation, where the only Python frame present is a DSL executor. - Whether such a fold is the author's own control flow or the DSL legitimately - truth-testing a wrapped value in glue code (config checks that never build - user-visible control flow) is a property of the live call stack, so it is - read off the consuming frame. Recording only author-frame folds keeps the - diagnostic free of false positives from DSL-internal probing. + Matches ``.func``-named ops plus :data:`_NON_DOT_FUNC_ENTRY_OPS` + (name-pattern, so new dialect function ops need no maintenance here). """ - try: - frame: "FrameType | None" = sys._getframe(1) - except Exception: # noqa: BLE001 - return False - from .common import is_dsl_internal_code # inline: break the import cycle + return op_name.endswith(".func") or op_name in _NON_DOT_FUNC_ENTRY_OPS - try: - while frame is not None: - if frame.f_globals.get("__name__") != __name__: - frame_is_internal = is_dsl_internal_code( - frame.f_code.co_filename, - frame.f_globals.get("__name__", "") or "", - ) - if not frame_is_internal: - return True - # The preprocessor RELOCATES user predicate expressions into - # DSL executor frames: `x in seq` -> in_ / compare chains, - # `a if c else b` -> ifExp_executor, `a and b` / `a or b` -> - # and_op / or_op. A fold consumed THERE is the user's - # control flow, not DSL-internal probing -- without this the - # witness for `if meta in (3,5,7):` / ternary / short-circuit - # predicates is dropped and a later promotion miscompiles - # silently. - return frame.f_code.co_name in _PRED_EXECUTOR_FRAMES - frame = frame.f_back - finally: - del frame - return False +def _is_module_boundary_op(op_name: str) -> bool: + """Return True if *op_name* is a module-like symbol-table container. -# DSL executor frames that evaluate USER predicate expressions on the -# preprocessor's behalf (see _fold_consumer_is_user_code). -# ``builtin_wrapper`` routes user calls to ``max``/``min``/``sum`` etc.; -# folds consumed inside it come from the USER's arguments (without -# a ``max(...)`` over watched values folded silently and a later -# promotion shipped the stale result). -_PRED_EXECUTOR_FRAMES: frozenset = frozenset( - {"ifExp_executor", "in_", "not_in", "and_op", "or_op", "not_op", "builtin_wrapper"} -) + The entry-block walk stops here: a module owns no SSA region for a + ``pyir.ref``, and walking past it risks recycled block wrappers. + """ + return op_name in _MODULE_OPS or op_name.endswith(".module") -def _wrapper_slots(value: Any) -> frozenset: - """All D1 slot keys a wrapper is connected to (own slot + provenance).""" - if not isinstance(value, _WatchedM): - return frozenset() - slots = value._origin_slots - if value._slot_key is not None: - slots = slots | {value._slot_key} - return slots +def _auto_promote_primitive(value: object) -> object | None: + """Promote a Python ``bool``/``int``/``float`` to the matching DSL type. + ``bool`` → ``Boolean``, ``int`` → ``Int32`` (``Int64`` for large), + ``float`` → ``Float32`` (via ``as_numeric`` from ``typing.py``). -def _record_fold_witness( - lhs: Any, rhs: Any, folded_value: Any, kind: str, force_consumer: bool = False -) -> None: - """Record that a slot-connected fold was consumed by plain CPython. - - No-op unless (a) inside staged CF, (b) at least one operand is connected - to a D1 slot, and (c) the consuming frame is author code (see - ``_fold_consumer_is_user_code``). Never raises -- the diagnostic is - deferred to ``_meta_promote_slot`` so slots that stay constexpr for the - whole trace (the supported meta-programming model) are never penalised. - - ``force_consumer`` skips gate (c): structural-integer consumption - (``__index__``) is dangerous no matter which frame performs it — the - canonical offender is ``range_constexpr``'s bound check, which calls - ``__index__`` from ``ast_helpers`` (a DSL frame) to fix a trace-time - unroll count that no later promotion can change. + Returns ``None`` on failure (import error, unsupported value). """ - slots = _wrapper_slots(lhs) | _wrapper_slots(rhs) - if not slots: - return - from .multi_stage_manager import is_inside_staged_cf - - if not is_inside_staged_cf(): - return - if not force_consumer and not _fold_consumer_is_user_code(): - return - from .diagnostics import find_user_source_location + try: + from .typing import as_numeric - filename, lineno, _col, _end_col = find_user_source_location() - record = { - "filename": filename, - "lineno": lineno, - "value": folded_value, - "kind": kind, - } - for slot in slots: - entries = _fold_witnesses.setdefault(slot, []) - if len(entries) < _MAX_FOLD_WITNESSES_PER_SLOT: - entries.append(record) - # A fold on a slot that is ALREADY runtime-varying can never be repaired - # by a later promotion (it already happened) -- deferring the diagnostic - # would leave this consumption silently stale (e.g. an alias captured - # before the region, structurally consumed after). Raise now. master - # tracks "already promoted" as membership in ``_slot_refs`` (set in - # ``_meta_promote_slot``), the bridge for the donor's ``_slot_stored``. - stale = [s for s in slots if s in _slot_refs] - if stale: - slot = stale[0] - name = slot[1] if isinstance(slot, tuple) and len(slot) > 1 else None - _raise_on_fold_witness(slot, folded_value, name, filename, lineno) - log().info( - "[fold-witness] %s consumed by CPython (%s) at %s:%s -> slots %s", - folded_value, - kind, - filename, - lineno, - list(slots), - ) + return as_numeric(value) + except (ImportError, TypeError, ValueError): + return None -# Op name -> Python operator; folds the shadow value and drives arith IR -# emission (delegation, !24500; consumed by ``_binop_ir``). -_BINOP_OPS: "dict[str, Callable[[Any, Any], Any]]" = { - "add": operator.add, - "sub": operator.sub, - "mul": operator.mul, - "truediv": operator.truediv, - "floordiv": operator.floordiv, - "mod": operator.mod, - "and": operator.and_, - "or": operator.or_, - "xor": operator.xor, - "lshift": operator.lshift, - "rshift": operator.rshift, -} - -# Op name -> Python operator; used to DELEGATE comparison (and arith) IR emission -# to the canonical ``ArithValue`` operator layer (which owns cmpi/cmpf predicate -# selection) -- consumed by ``_cmp_ir`` via ``_emit_arith_via_dsl`` / -# ``_emit_widened_via_numeric``. Superset of ``_BINOP_OPS`` (adds comparisons). -_DELEGATED_OP_FUNCS: "dict[str, Callable[[Any, Any], Any]]" = { - "add": operator.add, - "sub": operator.sub, - "mul": operator.mul, - "floordiv": operator.floordiv, - "truediv": operator.truediv, - "mod": operator.mod, - "and": operator.and_, - "or": operator.or_, - "xor": operator.xor, - "lshift": operator.lshift, - "rshift": operator.rshift, - "lt": operator.lt, - "le": operator.le, - "gt": operator.gt, - "ge": operator.ge, - "eq": operator.eq, - "ne": operator.ne, -} - - -def _needs_float_promotion(lhs_ir: "ir.Value", rhs: Any, op_name: str) -> bool: - """True when the op must use the promoting Numeric layer. - - Promote for: truediv on an integer lhs; a watched rhs (``ir.Value``) - whose type differs from the lhs; a plain float rhs on an integer lhs - (coercing it to the lhs type would truncate). A plain int rhs on a float - lhs coerces up exactly and stays on the ArithValue path. Pre-existing i1 - hole kept as-is: ``wb + 2`` still bakes the rhs to i1. - - (Rebase merge: the fold-witness ``_cmp_ir`` reuses this as its - float-widening discriminator -- the former ``_needs_float_widening`` had - byte-identical logic and was dropped.) - """ - lhs_is_float = isinstance(lhs_ir.type, ir.FloatType) - if op_name == "truediv" and not lhs_is_float: - return True - if isinstance(rhs, ir.Value): - return lhs_ir.type != rhs.type - return (not lhs_is_float) and isinstance(rhs, float) +# Lazily bound ``multi_stage_manager._is_staged_value`` (layering: that module +# imports from the pyir chain at module top, so the binding resolves on first use). +_IS_STAGED_VALUE_FN: "Callable[[object], bool] | None" = None -def _emit_arith_via_dsl(lhs_ir: "ir.Value", rhs_ir: "ir.Value", op_name: str) -> Any: - """Emit the arith/cmp op for *op_name* by DELEGATING to the DSL's canonical - ``ArithValue`` operator layer, rather than hand-picking arith ops here. +def _can_create_ref(value: object) -> bool: + """Return True if *value*'s type supports ``pyir.ref`` tracking. - ``ArithValue`` (``_mlir_helpers.arith``) already owns the mapping from a - Python operator to a concrete arith op -- ``cmpi``/``cmpf``, the - signed-vs-unsigned predicate choice, and so on. We wrap the - already-materialised operand SSA values in ``ArithValue`` and invoke the - operator. The operands are used verbatim -- the caller hands two - same-typed SSA values, so this never triggers type widening. + Only types with ``_pyir_ref_supported = True`` round-trip through + ``MutableValue`` (numerics, DSL pointers); Array/Tensor/TensorMap do not. - Returns the result ``ir.Value``, or ``NotImplemented`` for an op name not - in the table. Emission failures propagate so callers can apply their - existing NotImplemented / plain-fold fallback. + Pure (a plain ``getattr``): must NOT materialise ``ir_value()`` (it + would pin a constant and corrupt the sibling cache). """ - op_fn = _DELEGATED_OP_FUNCS.get(op_name) - if op_fn is None: - return NotImplemented + return getattr(type(value), "_pyir_ref_supported", False) - from .._mlir_helpers.arith import ArithValue - lhs_av = lhs_ir if isinstance(lhs_ir, ArithValue) else ArithValue(lhs_ir) - rhs_av = rhs_ir if isinstance(rhs_ir, ArithValue) else ArithValue(rhs_ir) - return op_fn(lhs_av, rhs_av) +def _is_scalar_ssa_carryable(value: object) -> bool: + """True if *value* is a staged non-compound leaf whose baked backing is exactly + one scalar int/float SSA value, so a single ``pyir.ref`` round-trips it.""" + global _IS_STAGED_VALUE_FN + try: + fn = _IS_STAGED_VALUE_FN + if fn is None: + from .multi_stage_manager import _is_staged_value + + fn = _IS_STAGED_VALUE_FN = _is_staged_value + if ( + fn(value) + and not isinstance(value, (bool, int, float, str, bytes, type)) + and not hasattr(value, "__extract_mlir_values__") + ): + raw = _raw_backing_ir_value(value) + if raw is not None and isinstance(raw.type, (ir.IntegerType, ir.FloatType)): + return True + except Exception: + return False + return False -# Reflected dunders: "wrap" boxes the plain lhs and runs the forward op; -# "swap" reuses the forward path (commutative mul). These fire only when the -# lhs type defers to the subclass -- ``2.0 / watched_int`` never dispatches -# here and stays unwatched. The other ops are absent on purpose: enabling -# one changes behavior and needs its own tests. -_REFLECTED_SPECS: "dict[str, str]" = { - "add": "wrap", - "sub": "wrap", - "mul": "swap", - "truediv": "wrap", -} +def _can_carry_leaf_ref(value: object) -> bool: + """True if a decomposition leaf can carry through one ``pyir.ref``: superset of + :func:`_can_create_ref` used only on decomposition / leaf-promotion paths.""" + return _can_create_ref(value) or _is_scalar_ssa_carryable(value) -def _generate_binop_dunders(cls: type) -> type: - """Generate the binary dunders from ``_BINOP_OPS`` / ``_REFLECTED_SPECS``. - Grep anchors: __add__ __sub__ __mul__ __truediv__ __floordiv__ __mod__ - __and__ __or__ __xor__ __lshift__ __rshift__ - __radd__ __rsub__ __rmul__ __rtruediv__ +def _is_vector_like(value: object) -> bool: + """Return True if *value* is a multi-element MLIR vector type. + + Vectors take the uniform place-cell rules; this predicate only selects + the generic ``pyir_read`` routing for loop-body write_args. """ + try: + ir_val = value.ir_value() # type: ignore[attr-defined] + return isinstance(ir_val.type, ir.VectorType) + except Exception: + return False - def _make_dunder(name: str, op: str, reflected_wrap: bool) -> Any: - if reflected_wrap: - def dunder(self: Any, other: Any) -> Any: - # Box the plain lhs so self's slot provenance threads through. - other_py = _unwrap(other) - if not isinstance(other_py, (bool, int, float)): - return NotImplemented - return _WatchedM(other_py)._binop_ir(self, op) +def _mlir_types_match(old_value: object, new_value: object) -> bool: + """True if both values carry the same MLIR type (read side-effect-free via + ``_raw_backing_ir_value``). Conservative: True when indeterminate.""" + old_raw = _raw_backing_ir_value(old_value) + new_raw = _raw_backing_ir_value(new_value) + if old_raw is None or new_raw is None: + return True + try: + return old_raw.type == new_raw.type + except Exception: + return True - else: - def dunder(self: Any, other: Any) -> Any: - return self._binop_ir(other, op) +def _types_match(old_value: object, new_value: object) -> bool: + """Full slot type identity (V-4): the MLIR half AND the wrapper-class half. + Conservative: True when either half is indeterminate.""" + if not _mlir_types_match(old_value, new_value): + return False + if type(old_value) is type(new_value): + return True + if _is_staged_value(old_value) and _is_staged_value(new_value): + return False + return True + + +def _is_memref_like(value: object) -> bool: + """Return True if *value* is a memref-backed descriptor (recomputable wherever + it dominates); declared via pyir's ``MemRefLikeTypeInterface``. Conservative: False.""" + try: + ir_type = value.type if isinstance(value, ir.Value) else value.ir_value().type # type: ignore[attr-defined] + except Exception: + return False + if isinstance(ir_type, ir.MemRefType): + return True + if pyir is None: + return False + try: + return bool(pyir.pyir_type_is_memref_backed(ir_type)) + except Exception: + return False - dunder.__name__ = name - dunder.__qualname__ = f"{cls.__qualname__}.{name}" - dunder.__doc__ = f"Staged ``{op}``; see ``_binop_ir``." - return dunder - for op in _BINOP_OPS: - setattr(cls, f"__{op}__", _make_dunder(f"__{op}__", op, False)) - for op, style in _REFLECTED_SPECS.items(): - setattr(cls, f"__r{op}__", _make_dunder(f"__r{op}__", op, style == "wrap")) - return cls +def _pyir_raise_memref_inregion_alloc_rebind(var: Any) -> None: + """Refuse loudly: a rebind under staged CF to a memref handle backed by an + IN-REGION allocation aliases ONE entry-hoisted buffer per iteration.""" + name = str(var) if var is not None else "" + raise DSLUserCodeError( + DiagId.MEMREF_INREGION_ALLOC_REBIND, + name=name, + region_kind=_pyir_enclosing_region_kind_at_ip(), + ) -def _emit_widened_via_numeric( - lhs_ir: "ir.Value", rhs_ir: "ir.Value", op_name: str -) -> Any: - """Emit *op_name* via the DSL Numeric layer, which applies int<->float type - widening (the ArithValue layer does not). Used only for the mixed / - truediv cases flagged by :func:`_needs_float_promotion`. +def _pyir_type_is_register_memref(ir_type: Any) -> bool: + """True iff *ir_type* is a REGISTER-backed memref-like handle (the + declared ``isRegisterBacked`` fact of pyir's MemRefLikeTypeInterface). + Conservative: False.""" + if pyir is None or not hasattr(pyir, "pyir_type_is_register_memref"): + return False + try: + return bool(pyir.pyir_type_is_register_memref(ir_type)) + except Exception: + return False - ``rhs_ir`` must carry its NATURAL type (int->i32, float->f32) so widening - sees the real operand kinds. Returns the result ``ir.Value`` (float for - arithmetic, i1 for comparisons) or ``NotImplemented``. - """ - op_fn = _DELEGATED_OP_FUNCS.get(op_name) - if op_fn is None: - return NotImplemented - from .typing import Numeric - from .._mlir_helpers.arith import ArithValue +# Iteration-private scratch admissions: register-space memref-handle rebinds +# the liveness gate admitted at the store choke. One record per admitted +# rebind store; the admitted cell threads as a carried buffer-identity phi +# (mem2reg promotes the slot into an scf iter_arg, the shape non-PyIR +# emits), and the trace-close sweep refuses any load whose consumers let a +# SUPERSEDED handle escape its window (see +# _pyir_verify_scratch_admissions). Consumed by the trace-close sweep. +_PYIR_SCRATCH_ADMISSIONS: "list[dict]" = [] + +# Last memref-handle store per cell: [ref, stored raw, store op] entries the +# store choke maintains (one per cell, latest wins). The auto-load stale-serve +# guard consults it to tell a wrapper holding the cell's CURRENT handle from a +# retained SUPERSEDED one. Trace-scoped. +_PYIR_MEMREF_LAST_STORE: "list[list]" = [] + + +def _pyir_note_memref_store(ref: "ir.Value", value: "ir.Value", store: Any) -> None: + """Record *store* as the latest memref-handle store into *ref*. Scan errors + propagate: a silently skipped match would leave a stale first entry shadowing + the update, making a superseded handle look CURRENT to the serve guard.""" + for entry in _PYIR_MEMREF_LAST_STORE: + if _same_ir_value(entry[0], ref): + entry[1] = value + entry[2] = store + return + _PYIR_MEMREF_LAST_STORE.append([ref, value, store]) + + +def _pyir_scratch_stale_serve_record(ref: Any, raw: "ir.Value | None") -> "dict | None": + """The admission record of *ref* when *raw* is a SUPERSEDED handle of that + admitted cell -- i.e. re-serving the wrapper through the cell would + observe a fresh generation's buffer instead of the buffer Python + retained. ``None`` when *ref* is not an admitted cell, *raw* is the + cell's current handle, or *raw* is a load taken after the last rebind. + Both record scans walk internally-built entries only; scan errors propagate + (an error swallowed to ``None`` here would serve the stale handle).""" + if raw is None or ref is None or not _PYIR_SCRATCH_ADMISSIONS: + return None + rec_found = None + for rec in _PYIR_SCRATCH_ADMISSIONS: + if _same_ir_value(rec["ref"], ref): + rec_found = rec + break + if rec_found is None: + return None + last_value = None + last_store = None + for entry in _PYIR_MEMREF_LAST_STORE: + if _same_ir_value(entry[0], ref): + last_value, last_store = entry[1], entry[2] + break + if last_value is None: + return None + if _same_ir_value(last_value, raw): + return None # the cell's current handle + try: + raw_owner = getattr(raw, "owner", None) + if raw_owner is not None and not isinstance(raw_owner, ir.Block): + raw_op = getattr(raw_owner, "operation", raw_owner) + if str(getattr(raw_op, "name", "")) == "pyir.load" and _same_ir_value( + raw_op.operands[0], ref + ): + # A load AFTER the last rebind is the current handle; a + # retained pre-rebind load is superseded. + if not _pyir_op_strictly_before(raw_op, last_store): + return None + except Exception: + pass + return rec_found - lhs_av = lhs_ir if isinstance(lhs_ir, ArithValue) else ArithValue(lhs_ir) - rhs_av = rhs_ir if isinstance(rhs_ir, ArithValue) else ArithValue(rhs_ir) - lhs_num = Numeric.from_mlir_type(lhs_av.type)(lhs_av) - rhs_num = Numeric.from_mlir_type(rhs_av.type)(rhs_av) - result_num = op_fn(lhs_num, rhs_num) - if result_num is NotImplemented: - return NotImplemented - return result_num.ir_value() +def _pyir_refuse_stale_memref_serve(value: Any, mv: Any) -> None: + """Refuse (loudly) serving *value* through its slot route *mv* when it is + a SUPERSEDED handle of an ADMITTED rebound memref cell: the cell now + carries a fresh generation, and under the one entry-hoisted buffer the + slot-routed read would observe that generation's overwrites instead of + the buffer Python retained. No-op for anything else.""" + if not _PYIR_SCRATCH_ADMISSIONS or mv is None: + return + ref = getattr(mv, "ref", None) + if ref is None: + return + raw = _raw_backing_ir_value(value) + if raw is None or not _is_memref_like(raw): + return + rec = _pyir_scratch_stale_serve_record(ref, raw) + if rec is None: + return + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.MEMREF_STALE_SCRATCH_CONSUMED, + filename=rec["filename"], + lineno=rec["lineno"], + name=rec["name"], + consumer_loc=( + f"{filename}:{lineno}" if filename is not None else "" + ), + ) -@_generate_binop_dunders -class _WatchedM: - """Wrap a Python primitive read inside staged CF for D1 tracking. - - The wrapper records the slot it came from so a later mutation can - rewrite baked ``arith.constant`` leaves to ``pyir.load %ref`` via - ``replaceAllUsesWith``. Arithmetic emits IR EAGERLY (returns a - derived wrapper carrying the result SSA); only LEAF constants are - tracked, downstream rewrite is handled by SSA edges when leaves are - replaced. - - **Crucial design point:** concrete instances subclass ``int`` or - ``float`` so that downstream code using ``isinstance(x, int)`` / - ``isinstance(x, float)`` sees them as primitives. Without this, - every consumer of an Mp-typed arg (``cutlass.Array(shape=...)``, - ``cute.printf``, MLIR builder shape params, etc.) would need a - band-aid unwrap -- D1's wrapping must be transparent to user code. - - Construction goes through ``_WatchedM(value, slot_key)`` which - dispatches to ``_WatchedInt`` (covers ``bool`` + ``int``) or - ``_WatchedFloat``. ``isinstance(x, _WatchedM)`` still works for - D1-internal checks because both concrete classes inherit from - ``_WatchedM``. See ``PYIR_DEV_GUIDE.md`` Pitfall 15. - """ - # Declared on the base class so type-checkers see them from - # ``ir_value``/``_pyir_auto_load_arg`` callers. Concrete subclass - # ``__new__`` populates each instance with the per-instance value. - _slot_key: Any - _cached_ir: Optional["ir.Value"] +def _pyir_op_transitive_consumer_loc(op: Any) -> "str | None": + """Source location of the first projection-closed consumer of *op*'s + results, or None when every transitive use is a dead projection chain. + The walk runs in C++: Python-side result access would run the dialect + value casters, which EMIT projection ops.""" + if pyir is None or not hasattr(pyir, "op_first_transitive_consumer_loc"): + return "" + try: + return pyir.op_first_transitive_consumer_loc(op) + except Exception: + return "" - # Slot provenance for DERIVED wrappers (``_slot_key is None``): the D1 - # slot keys of every tracked leaf this value was computed from, threaded - # through ``_binop_ir`` / ``_cmp_ir``. Lets a fold-witness on a derived - # value (e.g. ``if m > 15:`` truth-testing the derived ``_WatchedBool``) - # be attributed back to the leaf slot(s) whose promotion would make the - # fold stale. Class-level default keeps leaf construction unchanged. - _origin_slots: frozenset = frozenset() - def __new__(cls, value: Any = 0, slot_key: Any = None) -> "_WatchedM": - # Factory dispatch: _WatchedM(...) lands here and forwards to - # the right int/float-backed subclass. Subclass __new__ paths - # bypass this branch (cls is the subclass, not _WatchedM). - if cls is _WatchedM: - # Bool first because ``isinstance(True, int)`` is True. Keep - # bool routed through ``_WatchedBool`` so ``python_value`` - # returns ``True``/``False`` and downstream ``arith.const`` - # emits ``i1`` (not ``i32``). - if isinstance(value, bool): - return _WatchedBool.__new__(_WatchedBool, value, slot_key) - if isinstance(value, int): - return _WatchedInt.__new__(_WatchedInt, value, slot_key) - if isinstance(value, float): - return _WatchedFloat.__new__(_WatchedFloat, value, slot_key) - raise TypeError( - f"_WatchedM cannot wrap value of type {type(value).__name__}; " - "only bool/int/float are supported." +def _pyir_op_strictly_before(a: Any, b: Any) -> bool: + """True iff *a* executes strictly before *b* (ancestors ordered in their + deepest common block). Conservative: False.""" + if pyir is None or not hasattr(pyir, "op_strictly_before"): + return False + try: + return bool( + pyir.op_strictly_before( + getattr(a, "operation", a), getattr(b, "operation", b) ) - # Subclass __new__ already constructed the instance; nothing to do. - return super().__new__(cls) + ) + except Exception: + return False - @property - def python_value(self) -> Any: - """The bare Python value (for ``const_expr`` and similar). - Returns ``self`` cast to its int/float base via the subclass - override. Falls back to ``self`` for the base class (never - reached in practice). - """ - return self +def _pyir_classify_transitive_consumers( + op: Any, anchor: Any = None +) -> "tuple[int, str | None] | None": + """Consumer classification of *op*'s results under the projection closure + (CAPI fact): ``(mask, loc)`` with bit 1 = a consumer exists, bit 2 = a + consumer CAPTURES the value (stored as data / region-carried / terminator + / no effect anchored on the value; fail-closed), bit 4 = a non-capturing + consumer not strictly before *anchor*; ``loc`` names the first offending + consumer. None when the fact is unavailable (callers refuse).""" + if pyir is None or not hasattr(pyir, "op_classify_transitive_consumers"): + return None + try: + mask, loc = pyir.op_classify_transitive_consumers( + getattr(op, "operation", op), + None if anchor is None else getattr(anchor, "operation", anchor), + ) + return int(mask), loc + except Exception: + return None - def ir_value( - self, - *, - loc: "Optional[ir.Location]" = None, - ip: "Optional[ir.InsertionPoint]" = None, - ) -> "ir.Value": - """Emit (or reuse) the leaf ``arith.constant`` and record it. - - Cached across calls -- BUT the cache is only honoured when the - cached SSA still dominates the current insertion point. Without - that check, the same ``_WatchedInt`` instance read from N sibling - staged-CF regions (e.g. the eight ``scf.if`` bodies an unrolled - ``range_constexpr`` emits) would pin its only ``arith.constant`` - inside the FIRST region and leave N-1 dangling cross-region uses - -- the printer emits invalid IR ("use of undeclared SSA value - name") and the verifier reports the underlying region escape. - - Every materialisation is appended to ``_meta_uses[slot]`` so D1 - retroactive promotion can ``replaceAllUsesWith`` every baked - leaf, not just the first one. - """ - cached = self._cached_ir - if cached is not None and _cached_ir_value_dominates_ip(cached, ip): - return cached - const = _emit_constant_at_ip(self.python_value, loc=loc, ip=ip) - self._cached_ir = const - if self._slot_key is not None: - _meta_uses.setdefault(self._slot_key, []).append(const) - return const - # int / float / bool dunders are inherited from the int / float - # base class on the concrete subclass. ``__repr__`` is inherited - # too -- the printed form is the bare numeric, which is what we - # want (matches ``repr(int_value)``). +def _pyir_load_in_rebind_reexecution_scope(load_op: Any, region_op: Any) -> bool: + """True iff *load_op* sits inside *region_op* or inside a staged-CF + ancestor of it -- the scopes whose runtime re-execution re-runs the + rebind, so a handle captured there crosses an iteration boundary one + generation stale. Conservative: True on walk failure.""" + try: + cur = getattr(region_op, "operation", region_op) + while cur is not None: + name = getattr(cur, "name", None) + if name is None or _is_func_boundary_op(name): + return False + if name in _SCF_REGION_NAMES and _op_is_inside_op(load_op, cur): + return True + parent = cur.parent + cur = getattr(parent, "operation", parent) if parent is not None else None + return False + except Exception: + return True - # ----- arithmetic: emit IR eagerly, return derived wrappers ----------- - def _binop_ir(self, other: Any, op_name: str) -> Any: - """Fold the Python shadow and emit IR via the DSL operator layer. - - Same-typed operands go through ``ArithValue``; mixed / truediv combos - through the promoting ``Numeric`` layer (see - ``_needs_float_promotion``). Every failure except division-by-zero - returns ``NotImplemented`` so Python falls back to the ``int``/ - ``float`` base-class semantics. - """ - op_fn = _BINOP_OPS.get(op_name) - if op_fn is None: - return NotImplemented - try: - from .._mlir.dialects import arith as _arith - except ImportError: - return NotImplemented +def _pyir_admit_scratch_rebind( + value: "ir.Value", ref: "ir.Value", region_op: Any, var: Any +) -> bool: + """Liveness admission for an in-region register-memref rebind: the cell + threads as a loop-carried buffer-identity phi (mem2reg promotes the slot + into an scf iter_arg, exactly the shape non-PyIR emits), which is + observationally exact under the one entry-hoisted buffer UNLESS a + superseded handle stays reachable -- a CAPTURED old handle crosses an + iteration boundary one generation stale. Register space only (a + hardware fact: no async engine extends a register buffer's lifetime + beyond its handle). The past half is decided here (an in-region load of + the cell with a CAPTURING consumer refuses at the store); the witnessed + halves -- captures and consumers emitted AFTER this store -- are + enforced by the trace-close sweep over the recorded window.""" + if pyir is None or not hasattr(pyir, "op_classify_transitive_consumers"): + return False # no consumption facts -> the refusal stands + try: + if not _pyir_type_is_register_memref(value.type): + return False + region_norm = getattr(region_op, "operation", region_op) + for use in ref.uses: + user = getattr(use, "owner", None) + user_op = getattr(user, "operation", user) + if getattr(user_op, "name", "") != "pyir.load": + continue + if not _op_is_inside_op(user_op, region_norm): + continue + verdict = _pyir_classify_transitive_consumers(user_op) + if verdict is None or verdict[0] & 2: + return False # old handle captured in-region: live alias + except Exception: + return False + filename, lineno = _first_non_dsl_caller_location() + _PYIR_SCRATCH_ADMISSIONS.append( + { + "ref": ref, + "region": region_norm, + "stored_raw": value, + "name": str(var) if var is not None else "", + "filename": filename, + "lineno": lineno, + } + ) + return True - rhs_py = _unwrap(other) - if not isinstance(rhs_py, (bool, int, float)): - return NotImplemented +def _pyir_scratch_value_admitted(value: Any, region_op: Any) -> bool: + """True iff *value* is the stored handle of a scratch admission for + *region_op*: the region-close carry walk skips it (the admitted store + already went through the slot cell, which mem2reg promotes into the + carried phi -- a dedicated carry ref would double-carry; the admission + sweep owns the escape enforcement).""" + if not _PYIR_SCRATCH_ADMISSIONS or not isinstance(value, ir.Value): + return False + region_norm = getattr(region_op, "operation", region_op) + # A region-close RE-LOAD of the admitted slot names the same carried + # handle as the recorded store (the walk re-reads the cell): match the + # loaded-from ref alongside the raw stored handle. + loaded_ref = None + try: + def_op = _get_defining_operation(value) + if getattr(def_op, "name", None) == "pyir.load": + loaded_ref = def_op.operands[0] + except Exception: + loaded_ref = None + for rec in _PYIR_SCRATCH_ADMISSIONS: try: - result_py = op_fn(self.python_value, rhs_py) - except ZeroDivisionError: - # NotImplemented can't succeed here and would surface as a - # misleading TypeError for a float lhs. - raise + if rec["region"] != region_norm: + continue + if _same_ir_value(rec["stored_raw"], value): + return True + if loaded_ref is not None and _same_ir_value(rec["ref"], loaded_ref): + return True except Exception: - return NotImplemented - - try: - from .._mlir_helpers.arith import ArithValue - from .._mlir_helpers.arith import const as _arith_const - - def _av(v: Any) -> Any: - return v if isinstance(v, ArithValue) else ArithValue(v) + continue + return False - lhs_ir = self.ir_value() - if isinstance(other, _WatchedM): - rhs_ir = other.ir_value() - promote = _needs_float_promotion(lhs_ir, rhs_ir, op_name) - else: - # Natural type when promoting; else coerce to the lhs type. - promote = _needs_float_promotion(lhs_ir, rhs_py, op_name) - rhs_ir = ( - _arith_const(rhs_py) - if promote - else _arith_const(rhs_py, lhs_ir.type) - ) - if promote: - # Known waiver: promoted % emits remf while the shadow folds - # Python floor-mod (matches the non-PyIR runtime canon). - from .typing import Numeric +def _pyir_verify_scratch_admissions() -> None: + """Trace-close enforcement of the admitted rebinds: the cell itself + threads as a carried buffer-identity phi (a load of the cell always + names the CURRENT buffer), so the sweep polices only the escape routes + of a SUPERSEDED handle -- (a) a CAPTURE of any load of the cell inside a + scope whose runtime re-execution re-runs the rebind (the captured handle + crosses that scope's iteration boundary one generation stale), and (b) a + pre-rebind in-region load with a consumer not strictly before the next + rebind store, or with no positional window at all (e.g. a + while-condition read in the sibling region). Violations refuse loudly + at the consumer's source position. Consumes the admission list.""" + if not _PYIR_SCRATCH_ADMISSIONS: + return + records = list(_PYIR_SCRATCH_ADMISSIONS) + _PYIR_SCRATCH_ADMISSIONS.clear() + for rec in records: + region = rec["region"] + ref = rec["ref"] + + def _refuse(consumer_loc: "str | None") -> NoReturn: + raise DSLUserCodeError( + DiagId.MEMREF_STALE_SCRATCH_CONSUMED, + filename=rec["filename"], + lineno=rec["lineno"], + name=rec["name"], + consumer_loc=consumer_loc or "", + ) - result = op_fn( - Numeric.from_mlir_type(lhs_ir.type)(_av(lhs_ir)), - Numeric.from_mlir_type(rhs_ir.type)(_av(rhs_ir)), - ) - result_ir = ( - NotImplemented if result is NotImplemented else result.ir_value() - ) - elif op_name == "floordiv" and isinstance(lhs_ir.type, ir.FloatType): - # Legacy byte-compat: bare divf, no math.floor (ArithValue - # would add one; the shadow IS floored). - result_ir = _arith.divf(lhs_ir, rhs_ir) - else: - # ArithValue picks the arith op; i1 stays i1 (Numeric widens). - result_ir = op_fn(_av(lhs_ir), _av(rhs_ir)) - if result_ir is NotImplemented: - return NotImplemented + # The use walk feeds the refusal decision; the per-use probes are + # already conservative (defaulted getattrs, error-safe containment), + # so a raise here is an internal failure of the walk itself -- refuse + # rather than drop enforcement for this admitted rebind. + try: + loads: "list[Any]" = [] + stores_in: "list[Any]" = [] + for use in ref.uses: + user = getattr(use, "owner", None) + user_op = getattr(user, "operation", user) + op_name = getattr(user_op, "name", "") + if op_name == "pyir.load": + loads.append(user_op) + elif op_name == "pyir.store" and _op_is_inside_op(user_op, region): + stores_in.append(user_op) except Exception: - return NotImplemented + _refuse("") + for load_op in loads: + verdict = _pyir_classify_transitive_consumers(load_op) + if verdict is None: + _refuse("") + if verdict[0] == 0: + continue # dead projection chain (machinery seed): no read + anchor = None + if _op_is_inside_op(load_op, region): + after = [s for s in stores_in if _pyir_op_strictly_before(load_op, s)] + if after: + # Pre-rebind load: its window closes at the NEXT rebind. + anchor = after[0] + for s in after[1:]: + if _pyir_op_strictly_before(s, anchor): + anchor = s + elif not all(_pyir_op_strictly_before(s, load_op) for s in stores_in): + # A consumed load with no positional window w.r.t. some + # rebind store (e.g. a while-condition read in the + # sibling region): refuse. + _refuse(_pyir_op_transitive_consumer_loc(load_op)) + if anchor is not None: + verdict = _pyir_classify_transitive_consumers(load_op, anchor) + if verdict is None: + _refuse("") + mask, loc = verdict + if mask & 4: + _refuse(loc) # consumer escaped the pre-rebind window + if mask & 2 and _pyir_load_in_rebind_reexecution_scope(load_op, region): + _refuse(loc) # captured handle crosses an iteration boundary + + +def _pyir_ref_pointee_type_changed(ref: "ir.Value", value: object) -> bool: + """True iff *value*'s baked MLIR type differs from *ref*'s pointee (a ref is + single-typed; wrapper-class checks miss same-class rebinds). Conservative: False.""" + if ref is None: + return False + raw = _raw_backing_ir_value(value) + if raw is None: + return False + try: + return raw.type != ref.type.pointee + except Exception: + return False - derived = _WatchedM(result_py, slot_key=None) - derived._cached_ir = result_ir - # Thread slot provenance so a fold-witness on this derived value is - # attributed back to the leaf slot(s) whose promotion would make it - # stale (see ``_wrapper_slots`` / ``_record_fold_witness``). - derived._origin_slots = _wrapper_slots(self) | _wrapper_slots(other) - return derived - # Binary dunders are generated by ``_generate_binop_dunders``. +def _pyir_enclosing_region_op_at_ip() -> "ir.Operation | None": + """The innermost staged-CF region op enclosing the current insertion point, + or None outside any staged region. Conservative: None.""" + try: + block = ir.InsertionPoint.current.block + for _ in range(64): + op = block.owner + op = getattr(op, "operation", op) + name = getattr(op, "name", None) + if name in _SCF_REGION_NAMES: + return op + if ( + name is None + or _is_func_boundary_op(name) + or _is_module_boundary_op(name) + ): + break + parent_block = getattr(op, "block", None) + if parent_block is None: + break + block = parent_block + except Exception: + pass + return None - # ----- fold-witness channels: consumption that CPython forces to a bare - # Python value (no SSA edge), recorded so a later promotion raises - # PHASE_PREDICATE_FOLDED_STALE instead of silently shipping the stale fold. - def __bool__(self) -> bool: - """Truth-test: fold to the real Python truth value, but witness it. - - CPython requires a real ``bool`` here (``if w:``, ``while w:``, - ``not w``, ``and``/``or``), so the fold itself cannot be avoided. - Inside staged CF that fold decides a TRACE-TIME branch the IR never - sees -- sound only while the slot stays constexpr. Record a - fold-witness so a later promotion of the slot turns the silent - miscompile into ``DiagId.PHASE_PREDICATE_FOLDED_STALE`` (raised by - ``_meta_promote_slot``, which cites this fold's location). - """ - result = bool(self.python_value) - _record_fold_witness(self, None, result, "a Python `if`/`while`/`not` test") - return result +def _pyir_enclosing_region_kind_at_ip() -> str: + """Human-readable name of the innermost staged region enclosing the current + insertion point, for diagnostics; generic fallback when unresolvable.""" + op = _pyir_enclosing_region_op_at_ip() + name = getattr(op, "name", None) if op is not None else None + if name is None: + return "a for/while/if body" + return _SCF_REGION_NAMES.get(name, "a for/while/if body") - def __index__(self) -> int: - """Structural-integer consumption: loop bounds (``range_constexpr``), - list/tuple indices, shape parameters. The result fixes trace-time - CODE STRUCTURE (e.g. an unroll count) that no later promotion can - rewrite, so witness it even from DSL-internal frames - (``force_consumer``). CPython requires a real ``int``. - """ - pv = self.python_value - if not isinstance(pv, int): - raise TypeError( - f"'{type(self).__name__}' with non-int payload cannot be " - "interpreted as an integer" - ) - _record_fold_witness( - self, - None, - int(pv), - "a trace-time structural integer (constexpr loop bound / index / shape)", - force_consumer=True, - ) - return int(pv) - def __int__(self) -> int: - result = int(self.python_value) - _record_fold_witness(self, None, result, "an `int()` cast") - return result +def _pyir_raise_type_changed_in_region(var: Any, old_type: Any, new_type: Any) -> None: + """Refuse loudly: a type-changing rebind inside a staged region cannot be + lifted to a single-typed iter_arg/if-result and must not be silently skipped.""" + raise DSLUserCodeError( + DiagId.TYPE_CHANGED_INSIDE_REGION, + var=str(var) if var is not None else "", + old_type=str(old_type), + new_type=str(new_type), + region=_pyir_enclosing_region_kind_at_ip(), + ) - def __float__(self) -> float: - result = float(self.python_value) - _record_fold_witness(self, None, result, "a `float()` cast") - return result - def __format__(self, format_spec: str) -> str: - result = format(self.python_value, format_spec) - _record_fold_witness(self, None, result, "a string-formatting read") - return result - - # ----- comparisons: emit live cmpi/cmpf and return a derived wrapper so - # the predicate is a real SSA edge (lift-able to a dynamic while/ternary - # condition); fall back to a witnessed plain fold when IR cannot be made. - - def _cmp_ir(self, other: Any, op_name: str) -> Any: - # Staged DSL Numeric rhs (e.g. ``meta < i`` with a loop var): there - # is no Python fold — emit the cmp on the TRACKED leaf and return a - # staged Boolean so the predicate is live and promotion RAUW can - # rewrite the lhs. Without this the reflected Numeric compare - # bakes python_value as an UNTRACKED constant and a while on it - # never observes the body's writes (infinite scf.while). - from .typing import Boolean as _Boolean, Numeric as _Numeric - if isinstance(other, _Numeric): - try: - result_ir = _emit_arith_via_dsl( - self.ir_value(), - other.ir_value(), - op_name, - ) - if result_ir is NotImplemented: - return NotImplemented - return _Boolean(result_ir) - except Exception: - return NotImplemented +def _pyir_raise_rebind_duplicate_leaf(obj: Any) -> None: + """Refuse loudly: a whole-object rebind whose pre-region extract repeats + an SSA value has no faithful positional leaf pairing -- first-match would + silently cross-bind the duplicated positions.""" + raise DSLUserCodeError( + DiagId.WRAPPER_REBIND_DUPLICATE_LEAF, + cls=type(obj).__name__, + region=_pyir_enclosing_region_kind_at_ip(), + ) - # Non-primitive, non-Numeric rhs: defer to the rhs's reflected - # dunder via NotImplemented. For primitive-comparable operands we - # must NEVER return NotImplemented below this point: when both - # directions bail, CPython falls back to identity for ==/!= (two - # equal-valued wrappers compare unequal) and raises TypeError for - # orderings between two wrapper instances. - rhs_py = _unwrap(other) - if not isinstance(rhs_py, (bool, int, float)): - return NotImplemented - py_ops = { - "lt": lambda a, b: a < b, - "le": lambda a, b: a <= b, - "gt": lambda a, b: a > b, - "ge": lambda a, b: a >= b, - "eq": lambda a, b: a == b, - "ne": lambda a, b: a != b, - } - result_py = py_ops[op_name](self.python_value, rhs_py) +def _pyir_record_slot_template(place: Any, value: Any) -> None: + """F-TYPEID: advance the place row's wrapper-template half at a store or + publish choke; reads reconstruct the row from the recorded wrapper.""" + if place is None or value is None: + return + if isinstance(value, (_WatchedM, bool, int, float, str, bytes)): + return + _slot_templates[place] = value + + +def _declared_m2s_promotion_class(py_value: Any) -> "type | None": + """The ONE declared Python-scalar -> staged-class fact: bool->Boolean, + int->Int32 (widening to Int64 by value), float->Float32; ``None`` when no + value-preserving declared width exists.""" + from .typing import Boolean, Int32, Int64, Float32 + + if isinstance(py_value, bool): + return Boolean + if isinstance(py_value, int): + if -(2**31) <= py_value < 2**31: + return Int32 + if -(2**63) <= py_value < 2**63: + return Int64 + return None + if isinstance(py_value, float): + return Float32 + return None - try: - from .._mlir.dialects import arith as _arith # noqa: F401 - except ImportError: - # Plain-fold fallback on a slot-connected wrapper: the folded - # bool may feed control flow the meta table cannot see -- record - # a witness (no-ops when neither operand carries a slot). - _record_fold_witness( - self, other, result_py, "a comparison folded at trace time" - ) - return result_py - - # Disconnected wrapper (no slot, no derived SSA): nothing for - # promotion to rewrite — keep the plain Python fold. The RHS may - # still be slot-connected (e.g. ``_WatchedM(5) == watched_x``), so - # witness the fold for ITS slots (no-op when rhs carries none). - if self._slot_key is None and self._cached_ir is None: - _record_fold_witness( - None, other, result_py, "a comparison folded at trace time" - ) - return result_py +def _pyir_declared_promotion_template( + py_value: Any, staged_type: "type | None" = None +) -> Any: + """The declared M->S promotion identity: the staged write's own class, else + the scalar's declared promotion class (literal-backed, no IR emitted).""" + if staged_type is not None: try: - from .._mlir_helpers.arith import const as _arith_const - - lhs_ir = self.ir_value() - widen = _needs_float_promotion(lhs_ir, rhs_py, op_name) - if isinstance(other, _WatchedM): - rhs_ir = other.ir_value() - elif widen: - rhs_ir = _arith_const(rhs_py) # NATURAL type so widening sees it - else: - rhs_ir = _arith_const(rhs_py, lhs_ir.type) - if widen: - # int lhs vs a float rhs: widen to float so the compare - # matches Python. The same-type path would truncate the float - # rhs to the lhs integer type and flip results (e.g. ``w < 5.5`` - # with w==5 -> slt(5,5)==False instead of True). - result_ir = _emit_widened_via_numeric(lhs_ir, rhs_ir, op_name) - else: - result_ir = _emit_arith_via_dsl(lhs_ir, rhs_ir, op_name) - if result_ir is NotImplemented: - # Unreachable for the fixed set of comparison op names; route - # an unknown op through the plain-fold + witness path below - # rather than ever handing back a NotImplemented from a compare. - raise TypeError(f"unsupported comparison op {op_name!r}") - except Exception: - # IR emission failed (dead context, type mismatch such as an - # i1 leaf vs int rhs, mixed int/float) — keep exact Python - # comparison semantics rather than NotImplemented (see above). - # This is a PLAIN-FOLD consumption: the caller gets a bare - # ``bool`` with no live IR edge, so any control flow it feeds is - # invisible to promotion -- witness it. (The successful path - # below returns a derived wrapper with live ``cmpi``/``cmpf`` IR - # and is NOT witnessed; its truth-test, if any, is witnessed by - # ``__bool__`` instead.) - _record_fold_witness( - self, other, result_py, "a comparison folded at trace time" - ) - return result_py - - derived = _WatchedM(bool(result_py), slot_key=None) - derived._cached_ir = result_ir - derived._origin_slots = _wrapper_slots(self) | _wrapper_slots(other) - return derived - - def __lt__(self, other: Any) -> Any: - return self._cmp_ir(other, "lt") + return staged_type(py_value) + except (TypeError, ValueError, AttributeError): + pass + default_cls = _declared_m2s_promotion_class(py_value) + if default_cls is None: + return None + try: + return default_cls(py_value) + except (TypeError, ValueError, AttributeError): + return None - def __le__(self, other: Any) -> Any: - return self._cmp_ir(other, "le") - def __gt__(self, other: Any) -> Any: - return self._cmp_ir(other, "gt") +def _pyir_rebook_place_cell(mv: "MutableValue") -> None: + """After a type-transition re-mint, re-point the place's bare-ref row at the + live ref (the other place-keyed tables map to the MutableValue and follow).""" + place = getattr(mv, "_place", None) + if place is None or mv._ref is None: + return + if place in _slot_refs: + _slot_refs[place] = mv._ref + _pyir_record_slot_template(place, mv._value) - def __ge__(self, other: Any) -> Any: - return self._cmp_ir(other, "ge") - def __eq__(self, other: Any) -> Any: - return self._cmp_ir(other, "eq") +def _is_boolean_like(value: object) -> bool: + """Return True if *value* is a 1-bit integer (i1/Boolean) type. - def __ne__(self, other: Any) -> Any: - return self._cmp_ir(other, "ne") + Booleans take the uniform place-cell rules; this predicate only selects + the generic ``pyir_read`` routing for loop-body write_args. + """ + try: + ir_val = value.ir_value() # type: ignore[attr-defined] + return isinstance(ir_val.type, ir.IntegerType) and ir_val.type.width == 1 + except Exception: + return False -class _WatchedInt(_WatchedM, int): - """Concrete D1 wrapper backed by ``int``. +def _is_literal_backed(dsl_value: object) -> bool: + """Return True if *dsl_value* stores a Python scalar (not an ir.Value). - No ``__slots__`` -- CPython forbids non-empty ``__slots__`` on int - subclasses. We accept the small ``__dict__`` cost in exchange for - transparent ``isinstance(x, int)`` behavior throughout the DSL. + Literal-backed values mint a fresh ``arith.constant`` per ``ir_value()`` + (safe at block_begin); SSA-backed values must dominate the ref point. """ + return isinstance(getattr(dsl_value, "value", None), (bool, int, float)) - # Defining ``__eq__`` on ``_WatchedM`` implicitly sets ``__hash__ = - # None`` there; restore the base-class hash so watched values stay - # usable as dict keys / set members. - __hash__ = int.__hash__ - def __new__(cls, value: Any, slot_key: Any = None) -> "_WatchedInt": - inst = int.__new__(cls, value) - inst._slot_key = slot_key - inst._cached_ir = None - return inst +def _wrap_ir_like(template: Any, ir_val: "ir.Value") -> Any: + """The ONE reconstruction funnel: rebuild a DSL value from *ir_val* + mirroring *template* -- value-tree protocol, then shape/dtype ctor, then + single-arg ctor; a template with no viable protocol refuses loudly.""" + # A row with no wrapper template holds an unwrapped SSA binding: the loaded + # value IS the reconstruction (identity, not a guess). Exact-type test: + # dialect value-caster classes SUBCLASS ir.Value yet carry a wrapper + # identity and must take the strategies below. + if template is None or type(template) in ( + ir.Value, + ir.OpResult, + ir.BlockArgument, + ): + return ir_val + + # Strategy 1: value-tree reconstruct protocol (TensorSSA-family wrappers). + # Licensed ONLY by the full class-declared opt-in (BOTH protocol dunders): + # a half-opt-in masquerade falls through to ctor replay or the loud + # no-protocol refusal instead of an unlicensed reconstruct call. + if _implements_dynamic_expression(template): + new_from = getattr(template, "__new_from_mlir_values__", None) + if callable(new_from): + try: + return new_from([ir_val]) + except Exception: + pass # fall through - @property - def python_value(self) -> int: - # int.__int__ directly: ``int(self)`` would dispatch to the - # witness-recording ``_WatchedM.__int__`` override and recurse - # (it reads ``python_value`` right back). - return int.__int__(self) + # Strategy 2: replay constructor with shape/dtype metadata. + shape = getattr(template, "_shape", None) + if shape is not None: + dtype = getattr(template, "_dtype", None) + try: + return type(template)(ir_val, shape, dtype) + except Exception: + pass # fall through + # Strategy 3: simple single-arg constructor (scalar Numerics). + try: + return type(template)(ir_val) + except Exception: + raise DSLUserCodeError( + DiagId.RECONSTRUCT_NO_PROTOCOL, cls=type(template).__name__ + ) from None -class _WatchedBool(_WatchedM, int): - """Concrete D1 wrapper for ``bool`` values. - - ``bool`` is final in CPython so we cannot subclass it; subclass - ``int`` (which ``bool`` itself subclasses) and override - ``python_value`` to return a proper ``bool``. This keeps - ``arith.const`` on ``python_value`` emitting ``i1``, while - ``isinstance(x, int)`` still passes downstream. Note: - ``isinstance(x, bool)`` returns ``False`` -- consumers that need - bool-specific behavior must coerce via ``python_value`` or call - ``arith.const`` which routes ``_WatchedM`` through ``ir_value()``. - """ - __hash__ = int.__hash__ +def _make_poison_like(dsl_value: Any, ir_val: "ir.Value") -> Any: + """Create a ``ub.poison`` value of the same type as *dsl_value*. - def __new__(cls, value: Any, slot_key: Any = None) -> "_WatchedBool": - inst = int.__new__(cls, int(value)) - inst._slot_key = slot_key - inst._cached_ir = None - return inst + Case D-fallback ref init: the caller MUST ``store`` before any ``load``; + a pre-store read surfaces the poison instead of a silent zero. - @property - def python_value(self) -> bool: - return bool(int.__int__(self)) # bypass _WatchedM.__int__ (recursion) + Stamps the first non-DSL caller (file, line) on the op for the poison-read + diagnostic; wraps via :func:`_wrap_ir_like`; zero-constant without ``ub``. + """ + if ub is None: + return _make_zero_like(dsl_value, ir_val) + poison_ir = ub.PoisonOp(ir_val.type).result + _record_poison_source(poison_ir) + return _wrap_ir_like(dsl_value, poison_ir) + + +def _make_zero_like(dsl_value: Any, ir_val: "ir.Value") -> Any: + """DEFINED zero/false placeholder init of *dsl_value*'s type: only dead bypass + paths ever observe it, so it must not be a surviving ``ub.poison``.""" + zero_ssa = _make_raw_placeholder_init(ir_val.type) + if zero_ssa is not None: + return _wrap_ir_like(dsl_value, zero_ssa) + # Non-scalar / unrecognised type: fall back to a type-correct ``ub.poison``; + # the store-after-def still dominates every real read. + if ub is not None: + poison_ir = ub.PoisonOp(ir_val.type).result + _record_poison_source(poison_ir) + return _wrap_ir_like(dsl_value, poison_ir) + # No ``ub`` and not a recognised scalar: last-resort literal wrapper. + return type(dsl_value)(0) + + +def _pyir_type_is_scalar(ir_type: "ir.Type") -> bool: + """True for scalar MLIR types (int/index/standard float): safe to materialize + in Python; any other type may have a dialect value caster that emits ops.""" + try: + return isinstance( + ir_type, + ( + ir.IntegerType, + ir.IndexType, + ir.F16Type, + ir.F32Type, + ir.F64Type, + ir.BF16Type, + ), + ) + except Exception: + return False -class _WatchedFloat(_WatchedM, float): - """Concrete D1 wrapper backed by ``float``. Same ``__slots__`` rule - as :class:`_WatchedInt`.""" +def _make_raw_placeholder_init(ir_type: "ir.Type") -> "ir.Value | None": + """Emit a stamped zero/false ``arith.constant`` stand-in init for a SCALAR + *ir_type*, as a bare ``ir.Value`` (never DSL-wrapped); ``None`` for non-scalar.""" + from .._mlir.dialects import arith as _arith - __hash__ = float.__hash__ + zero_ssa: "ir.Value | None" = None + try: + if isinstance(ir_type, (ir.IntegerType, ir.IndexType)): + zero_ssa = _arith.constant(ir_type, 0) + elif isinstance(ir_type, (ir.F16Type, ir.F32Type, ir.F64Type, ir.BF16Type)): + zero_ssa = _arith.constant(ir_type, 0.0) + except Exception: + zero_ssa = None + if zero_ssa is not None and isinstance(zero_ssa, ir.Value): + _record_placeholder_init_source(zero_ssa) + return zero_ssa + return None - def __new__(cls, value: Any, slot_key: Any = None) -> "_WatchedFloat": - inst = float.__new__(cls, value) - inst._slot_key = slot_key - inst._cached_ir = None - return inst - @property - def python_value(self) -> float: - # float.__float__ directly: ``float(self)`` would dispatch to the - # witness-recording ``_WatchedM.__float__`` override and recurse. - return float.__float__(self) +def _mint_write_only_placeholder_cell( + dsl_value: Any, ir_type: "ir.Type", entry_block: "ir.Block" +) -> "MutableValue | None": + """Mint a write-only placeholder cell for a non-scalar type entirely in C++ + (the poison init never materializes in Python); ``None`` when unavailable.""" + if pyir is None or not hasattr(pyir, "mint_write_only_placeholder_cell"): + return None + src_file, src_line = _first_non_dsl_caller_location() + try: + ref = pyir.mint_write_only_placeholder_cell( + ir_type, entry_block, src_file or "", src_line or 0 + ) + except Exception: + return None + if not isinstance(ref, ir.Value): + return None + # The C++ helper stamped a `ub.poison` placeholder init that Python never + # touches; count it so the verify-boundary fast path stays sound. + _count_poison_emitted() + mv = MutableValue(dsl_value) + mv._ref = ref + mv._ref_context_id = id(ir.Context.current) + return mv -def _replace_value_uses(old_val: "ir.Value", new_val: "ir.Value") -> bool: - """Best-effort ``replaceAllUsesWith`` across MLIR Python binding versions.""" - for method_name in ("replace_all_uses_with", "replaceAllUsesWith"): - fn = getattr(old_val, method_name, None) - if fn is not None: - try: - fn(new_val) - return True - except Exception: - continue - return False +def _mint_ref_with_raw_init(dsl_value: Any, init_ir: "ir.Value") -> "MutableValue": + """Build a ``MutableValue`` whose ref is minted from the raw stand-in + *init_ir*; the wrapper template stays the REAL value so loads reconstruct.""" + mv = MutableValue(dsl_value) + mv._ref = pyir.ref(init_ir) + mv._ref_context_id = id(ir.Context.current) + return mv -def _promoted_type_name(py_value: Any) -> str: - """DSL type name a Python primitive promotes to (for diagnostics). +# Bound on the lazy caller-location climb: a pathological stack must not turn +# a diagnostic annotation into a linear-cost walk. +_CALLER_LOCATION_MAX_CLIMB = 256 - Thin wrapper over the canonical mapping ``Numeric.from_python`` (single - mechanism for Python-primitive -> DSL-type deduction); falls back to the - raw Python type name for non-primitives so diagnostics never raise. - """ - from .typing import Numeric +def _first_non_dsl_caller_location() -> "tuple[str | None, int | None]": + """(filename, lineno) of the first non-DSL caller frame, for diagnostics only; + a bounded ``sys._getframe`` climb keyed on module top-level package name.""" + dsl_top_level = __name__.split(".", 1)[0] try: - return Numeric.from_python(py_value).__name__ - except Exception: # noqa: BLE001 -- diagnostics must not raise - return type(py_value).__name__ + # Start two frames up: skip this helper and its recording caller, + # matching the frame selection the renderers rely on. + frame: "types.FrameType | None" = sys._getframe(2) + except ValueError: + return None, None + for _ in range(_CALLER_LOCATION_MAX_CLIMB): + if frame is None: + break + mod_name = frame.f_globals.get("__name__", "") + if not (mod_name == dsl_top_level or mod_name.startswith(dsl_top_level + ".")): + return frame.f_code.co_filename, frame.f_lineno + frame = frame.f_back + return None, None -def _raise_on_fold_witness( - slot_key: Any, - initial_py_value: Any, - target_name: "str | None", - filename: "str | None", - lineno: "int | None", -) -> None: - """Refuse to promote a slot whose fold already decided CPython control flow. - - A recorded fold-witness means a trace-time branch was taken from this - slot's constexpr value; promoting the slot now would make its reads - runtime loads while that branch stays hard-wired -- a guaranteed silent - miscompile. Raise ``DiagId.PHASE_PREDICATE_FOLDED_STALE`` citing BOTH - the fold site (from the witness) and the mutation site (this promotion's - ``filename``/``lineno``, which also anchor the rendered code frame). - """ - witnesses = _fold_witnesses.get(slot_key) - if not witnesses: - return - from .common import DSLUserCodeError - from .diagnostics import DiagId, find_user_source_location +# Poison / placeholder inits stamped since the last verify boundary; the +# used-poison scan skips the module walk when zero. +_POISON_EMITTED: int = 0 - if filename is None and lineno is None: - filename, lineno, _col, _end_col = find_user_source_location() - first = witnesses[0] - fold_location = ( - f"{first['filename']}:{first['lineno']}" - if first["filename"] is not None - else "a location inside a plain-Python helper" - ) - if len(witnesses) > 1: - fold_location += f" (and {len(witnesses) - 1} more place(s))" - mut_location = ( - f"{filename}:{lineno}" if filename is not None else "the highlighted line" - ) - var = target_name if target_name is not None else "this value" - raise DSLUserCodeError( - DiagId.PHASE_PREDICATE_FOLDED_STALE, - filename=filename, - lineno=lineno, - var=var, - value=repr(initial_py_value), - type=_promoted_type_name(initial_py_value), - fold_location=fold_location, - fold_kind=first["kind"], - mut_location=mut_location, - ) +def _count_poison_emitted() -> None: + global _POISON_EMITTED + _POISON_EMITTED += 1 -def _exit_function_trace() -> None: - """Clear D1 per-trace state. Called from ``_jit_scope.finally``.""" - _meta_uses.clear() - _slot_refs.clear() - _slot_first_def_inside_cf.clear() - _slot_first_def_depth.clear() - _slot_first_def_depth_any.clear() - _slot_first_def_block.clear() - _slot_mvs.clear() - _fold_witnesses.clear() # P6: per-trace fold-witness table - _slot_pending_store.clear() # pfound: per-region syntactic write-set - _pyir_fn_id_stack.clear() - _pinned_owners.clear() # release owners we kept alive for slot-key identity +def _record_poison_source(poison_value: "ir.Value") -> None: + """Stamp the first non-DSL caller location onto the fresh ``ub.poison`` op so + the end-of-trace catcher can render a diagnostic without MLIR loc info.""" + _count_poison_emitted() + src_file, src_line = _first_non_dsl_caller_location() + if src_file is None: + return + try: + poison_op = poison_value.owner + poison_op.attributes["pyir.poison_src_file"] = ir.StringAttr.get(src_file) + poison_op.attributes["pyir.poison_src_line"] = ir.IntegerAttr.get( + ir.IntegerType.get_signless(32), src_line + ) + except (AttributeError, TypeError): + pass -def _watched_to_dsl(watched: "_WatchedM") -> Any: - """Convert a ``_WatchedM`` wrapper into the matching DSL Numeric. +def _record_placeholder_init_source(placeholder_value: "ir.Value") -> None: + """Mark a placeholder-init constant (unconditional ``pyir.placeholder_init`` + attr + best-effort source attrs) so an uncovered read of it faults loudly.""" + _count_poison_emitted() + try: + ph_op = placeholder_value.owner + ph_op.attributes["pyir.placeholder_init"] = ir.UnitAttr.get() + except (AttributeError, TypeError): + return + src_file, src_line = _first_non_dsl_caller_location() + if src_file is None: + return + try: + ph_op.attributes["pyir.placeholder_init_src_file"] = ir.StringAttr.get(src_file) + ph_op.attributes["pyir.placeholder_init_src_line"] = ir.IntegerAttr.get( + ir.IntegerType.get_signless(32), src_line + ) + except (AttributeError, TypeError): + pass - Bakes the leaf ``arith.constant`` and records it in ``_meta_uses`` - (if leaf), or returns the cached IR (if derived from arithmetic). - Wraps the resulting ``ir.Value`` in Boolean / Int32 / Float32 based - on the Python type of the wrapped value, so downstream DSL APIs - (``cute.printf`` etc.) see a familiar Numeric. - """ - ir_val = watched.ir_value() - py = watched.python_value - from .typing import Boolean, Int32, Float32 - if isinstance(py, bool): - return Boolean(ir_val) - if isinstance(py, int): - return Int32(ir_val) - if isinstance(py, float): - return Float32(ir_val) - return ir_val +def _get_defining_operation(ir_val: ir.Value) -> ir.Operation: + """Return the ``ir.Operation`` that defines an ``OpResult``.""" + owner = ir_val.owner + return getattr(owner, "operation", owner) -def _describe_value_origin(raw: Any) -> str: - """Describe where an MLIR value was defined (best-effort). +def _get_function_entry_block() -> ir.Block | None: + """Walk up from the current insertion point to find the enclosing + function-like op's entry block. - Returns a string like ``"inside a for-loop body"`` or ``""``. + Returns the entry block, or ``None``. Stops at the nearest function + boundary and never walks past a module container (recycled-wrapper hazard). """ try: - owner = raw.owner + block = ir.InsertionPoint.current.block except Exception: - return "" + return None - # Find the parent Operation that owns the region containing this value. - # BlockArgument: owner is Block -> Block.region -> Region.owner - # OpResult: owner is Operation -> Operation.block -> Block.region -> Region.owner - try: - if isinstance(owner, ir.Block): - parent_name = owner.region.owner.name - elif isinstance(owner, ir.Operation): - parent_name = owner.block.region.owner.name - else: - return "" - except Exception: - return "" + # Bound the walk by IR nesting depth, not an id() visited-set: recycled + # wrapper ids / a non-None recycled op.block could loop into a SIGSEGV. + _MAX_NESTING = 256 - desc = _SCF_REGION_NAMES.get(parent_name, f"a `{parent_name}` region") - return f"inside {desc}" + for _ in range(_MAX_NESTING): + if block is None: + break + try: + parent = block.owner # Block → Python dialect op + except Exception: + break + op = getattr(parent, "operation", parent) # → ir.Operation + try: + op_name = str(op.name) + except Exception: + break -# --------------------------------------------------------------------------- -# Ref placement helpers -# --------------------------------------------------------------------------- + if _is_func_boundary_op(op_name): + return op.regions[0].blocks[0] -_FUNC_OPS = frozenset(("func.func", "gpu.func", "cuda.kernel", "llvm.func")) + # A module container is the symbol-table root (no enclosing function); stop + # rather than walk into the binding's recycled module-block aliasing. + if _is_module_boundary_op(op_name): + break + try: + block = op.block # ir.Operation → parent Block + except Exception: + break -def _auto_promote_primitive(value: object) -> object | None: - """Promote a Python ``bool``/``int``/``float`` to the matching DSL type. + return None - Uses ``as_numeric`` from ``typing.py``: - ``bool`` → ``Boolean``, ``int`` → ``Int32`` (``Int64`` for large), - ``float`` → ``Float32``. - Returns ``None`` on failure (import error, unsupported value). - """ +def _pyir_recorded_birth_block( + target_name: "str | None", owner: Any, slot_name: Any +) -> "ir.Block | None": + """The place's recorded F-BIRTHPOS block (``None`` when unrecorded).""" try: - from .typing import as_numeric - - return as_numeric(value) - except (ImportError, TypeError, ValueError): + slot = _make_slot_key(target_name, owner, slot_name) + except Exception: return None + if slot is None: + return None + return _slot_first_def_block.get(slot) -def _can_create_ref(value: object) -> bool: - """Return True if *value*'s type supports ``pyir.ref`` tracking. - - Only types with ``_pyir_ref_supported = True`` can round-trip through - ``MutableValue``: ``Type(ir_value)`` must be a valid single-arg - constructor. Numeric types (Int32, Float32, Boolean, CuTe Pointer) - set this. Non-scalar types (Array, Tensor, TensorMap) do not. - """ - return getattr(type(value), "_pyir_ref_supported", False) +def _mint_region_born_cell( + dsl_value: Any, birth_block: "ir.Block", entry_block: "ir.Block" +) -> "MutableValue | None": + """Mint the cell for a region-born place: undefined-placeholder entry init + (dominance device) + the seed stored at the recorded birth block + (invariant E). + ``None`` when no placeholder init is constructible (legacy placement + applies); refuses when the seed has no position inside the birth block.""" + # A memref-backed descriptor is a recomputable value handle with no entry + # cell (Case D's declared arm): seed semantics do not apply. + if _is_memref_like(dsl_value): + return None + # Locate the seed at its binding position: a literal re-materialises inside + # the birth block; an SSA seed keeps its defining position. + seed_raw = _raw_backing_ir_value(dsl_value) + if seed_raw is None: + if _is_literal_backed(dsl_value): + with ir.InsertionPoint.at_block_begin(birth_block): + seed_raw = dsl_value.ir_value() + else: + seed_raw = dsl_value.ir_value() + # Seed placement: after a same-block def; at block begin for a dominating + # outer def; no position exists otherwise (LAW 3 fabricates none). + seed_owner = seed_raw.owner + if isinstance(seed_owner, ir.Block): + seed_def_op = None + seed_in_birth = seed_owner == birth_block + else: + seed_def_op = getattr(seed_owner, "operation", seed_owner) + seed_in_birth = seed_def_op.block == birth_block + if not seed_in_birth and not pyir.value_dominates_ip(seed_raw, birth_block, None): + src_file, src_line = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.BODY_BORN_SEED_UNPLACEABLE, + filename=src_file, + lineno=src_line, + var=type(dsl_value).__name__, + ) + mv: "MutableValue | None" + with ir.InsertionPoint.at_block_begin(entry_block): + init_ir = _make_raw_placeholder_init(seed_raw.type) + if init_ir is not None: + mv = _mint_ref_with_raw_init(dsl_value, init_ir) + else: + mv = _mint_write_only_placeholder_cell( + dsl_value, seed_raw.type, entry_block + ) + if mv is None: + return None + if seed_in_birth and seed_def_op is not None: + seed_ip = ir.InsertionPoint.after(seed_def_op) + else: + seed_ip = ir.InsertionPoint.at_block_begin(birth_block) + # Init-like raw emission (the mint is not a store event): the cell's + # template already IS *dsl_value*, so no wrapper facts need advancing. + with seed_ip: + _pyir_emit_store(seed_raw, mv._ref, choke="region-born birth seed") + return mv -def _is_vector_like(value: object) -> bool: - """Return True if *value* is a multi-element MLIR vector type. - Vectors are ``_pyir_ref_supported`` (they round-trip via - ``Vector(ir_value)``), but creating first-def refs for them can - cause issues when they are subsequently stored into dicts or passed - across MLIR region boundaries. The reassignment path handles them - correctly when the old value already exists. - """ - try: - ir_val = value.ir_value() # type: ignore[attr-defined] - return isinstance(ir_val.type, ir.VectorType) - except Exception: - return False +def _create_ref( + dsl_value: Any, birth_block: "ir.Block | None" = None +) -> "MutableValue": + """Create a ``MutableValue`` + ``pyir.ref`` with correct placement. + *birth_block* is the place's recorded first-def block (F-BIRTHPOS); a + region-born place takes an undefined-placeholder entry init plus a real + seed store at that block, so the binding executes where Python executes it. -def _is_memref_like(value: object) -> bool: - """Return True if *value* is backed by a memref/tensor-descriptor SSA value - (e.g. a ``cute._Tensor``). - - Such a value is a *descriptor* (pointer + layout), recomputable wherever it - dominates. Inside a nested region the normal Case-D path would put a - ``ub.poison``-init ref at the function entry block; for a memref that poison - leaks across sibling regions (the P89 regression). These values dominate the - current IP, so the ref can simply be placed there instead. - - Detection is declarative: the value's type opts in via the class attribute - ``_pyir_memref_backed`` (``cute._Tensor`` sets it). ``base_dsl`` stays - decoupled from the cute dialect -- no import of, or string-matching against, - cute's MLIR type names. Builtin ``memref`` values are also accepted via a - structural ``isinstance`` check. - """ - if getattr(type(value), "_pyir_memref_backed", False): - return True - try: - return isinstance(value.ir_value().type, ir.MemRefType) # type: ignore[attr-defined] - except Exception: - return False + Placement strategy (4 cases, most optimal to most conservative): + **Case A** — Literal-backed: function entry block begin (a fresh + ``arith.constant`` per ``ir_value()``, dominance guaranteed). -def _mlir_type_or_none(value: object) -> "ir.Type | None": - """Return *value*'s MLIR type, or ``None`` if it has no SSA backing.""" - try: - return value.ir_value().type # type: ignore[attr-defined] - except Exception: - return None + **Case B** — Entry-block argument: entry block begin (block args + dominate their whole block). + **Case C** — Entry-block OpResult: immediately after the defining op. -def _staged_type_changed(old_value: object, new_value: object) -> bool: - """Return True when a same-name reassignment changes the MLIR type. + **Case D** — Everything else: current insertion point (the value + dominates it — Python is using it here). - A ``pyir.ref`` has an element type that is fixed at creation. When a - same-name local is reassigned to a value of a DIFFERENT MLIR type - inside staged CF (e.g. ``v = v.to(other_dtype)`` or a vector recast - ``vector<8xf4> -> vector<4xi8>``), the ref created for the old value - can no longer hold the new one: ``pyir.store``/``pyir.load`` would - round-trip the OLD type and downstream uses read a stale-typed value. - Detecting the change lets the reassignment path drop the old ref and - create a fresh one typed to the new value. - - Both operands must have an MLIR type for the comparison to be - meaningful; if either is unavailable we conservatively report "no - change" so the legacy ref-reuse path is preserved. - """ - old_ty = _mlir_type_or_none(old_value) - new_ty = _mlir_type_or_none(new_value) - if old_ty is None or new_ty is None: - return False - return old_ty != new_ty - - -def _is_boolean_like(value: object) -> bool: - """Return True if *value* is a 1-bit integer (i1/Boolean) type. - - Boolean first-defs (e.g. ``cond = i > 2``) are almost always used - once as an ``scf.if`` condition and never modified. Creating an - eager first-def ``pyir.ref`` for them produces a dead ref that the - C++ ``PYIRToSCF`` pass turns into a spurious iter_arg with a - duplicate ``scf.yield``, breaking MLIR verification. - - The reassignment path still handles booleans correctly when an - existing ``_mutable_ref`` is already attached (from being modified - inside control flow). - """ - try: - ir_val = value.ir_value() # type: ignore[attr-defined] - return isinstance(ir_val.type, ir.IntegerType) and ir_val.type.width == 1 - except Exception: - return False - - -def _is_literal_backed(dsl_value: object) -> bool: - """Return True if *dsl_value* stores a Python scalar (not an ir.Value). - - Literal-backed values produce a fresh ``arith.constant`` on every - ``ir_value()`` call, so ``pyir.ref`` can be placed at block_begin - without dominance issues. SSA-backed values return an existing - ``ir.Value`` which must dominate the ref insertion point. - """ - return isinstance(getattr(dsl_value, "value", None), (bool, int, float)) - - -# Count of ``ub.poison`` ops built by ``_make_poison_like`` — the ONLY -# Python-layer producer of ``ub.poison`` — since the last verify boundary. -# ``_verify_no_used_poison`` walks the whole module looking for used poison; -# while this counter is zero no module can contain one, so the walk (which -# runs on EVERY compile, pyir or classic) is skipped. The verifier consumes -# the counter at its boundary so later poison-free compiles in the same -# process keep the fast path (traces are sequential in-process). -_POISON_EMITTED: int = 0 - - -def _make_poison_like(dsl_value: Any, ir_val: "ir.Value") -> Any: - """Create a ``ub.poison`` value of the same type as *dsl_value*. - - Used by ``_create_ref`` Case D-fallback when the original SSA value - does not dominate the current insertion point. The poison signals - deferred undefined behaviour: the caller MUST ``store`` a real value - before any ``load``. If that contract is ever failed, the poison - makes the bug visible instead of silently returning zero. - - *ir_val* is the already-computed ``dsl_value.ir_value()`` — used - only for its type. - - Records the (filename, lineno) of the first non-DSL caller as MLIR - attributes on the poison op so the end-of-trace poison-read catcher - can render an actionable diagnostic even when the MLIR ``loc()`` was - stripped. - - When the ``ub`` dialect is unavailable, falls back to a zero constant - of the matching type. - """ - if ub is None: - ir_type = ir_val.type - if ir_type == ir.IntegerType.get_signless(1): - return type(dsl_value)(False) - if isinstance(ir_type, ir.IntegerType): - return type(dsl_value)(0) - return type(dsl_value)(0.0) - global _POISON_EMITTED - _POISON_EMITTED += 1 - poison_ir = ub.PoisonOp(ir_val.type).result - _record_poison_source(poison_ir) - return type(dsl_value)(poison_ir) - - -def _record_poison_source(poison_value: "ir.Value") -> None: - """Stash the first pure Python caller location as MLIR attributes - on the freshly-created ``ub.poison`` op, so the end-of-trace catcher - can render an actionable diagnostic regardless of MLIR loc info. - - "DSL-internal" is determined by the top-level component of each - frame's module ``__name__``: any frame whose module shares this - file's top-level package is skipped. The import system sets - ``__name__`` from how the module was imported, not from a filesystem - path — so this works identically under pip install, build symlinks, - in-source dev, and any CI runner without any hardcoded paths. - Tests run as ``__main__`` and user kernels as their own module names; - neither shares the DSL's top-level, so both are correctly reported. - """ - dsl_top_level = __name__.split(".", 1)[0] - for frame_info in inspect.stack()[1:]: - mod_name = frame_info.frame.f_globals.get("__name__", "") - if mod_name == dsl_top_level or mod_name.startswith(dsl_top_level + "."): - continue - try: - poison_op = poison_value.owner - poison_op.attributes["pyir.poison_src_file"] = ir.StringAttr.get( - frame_info.filename - ) - poison_op.attributes["pyir.poison_src_line"] = ir.IntegerAttr.get( - ir.IntegerType.get_signless(32), frame_info.lineno - ) - except (AttributeError, TypeError): - pass - return - - -def _get_defining_operation(ir_val: ir.Value) -> ir.Operation: - """Return the ``ir.Operation`` that defines an ``OpResult``.""" - owner = ir_val.owner - return getattr(owner, "operation", owner) - - -def _get_function_entry_block() -> ir.Block | None: - """Walk up from the current insertion point to find the enclosing - function-like op's entry block. - - Returns the entry block, or ``None`` if no function is found. - Naturally respects ``IsolatedFromAbove`` — stops at the nearest - function boundary (``func.func``, ``gpu.func``, ``cuda.kernel``). - """ - try: - block = ir.InsertionPoint.current.block - except Exception: - return None - - while block is not None: - try: - parent = block.owner # Block → Python dialect op - except Exception: - break - op = getattr(parent, "operation", parent) # → ir.Operation - - if str(op.name) in _FUNC_OPS: - return op.regions[0].blocks[0] - - try: - block = op.block # ir.Operation → parent Block - except Exception: - break - - return None - - -def _create_ref(dsl_value: Any) -> "MutableValue": - """Create a ``MutableValue`` + ``pyir.ref`` with correct placement. - - Placement strategy (4 cases, most optimal to most conservative): - - **Case A** — Literal-backed (``self.value`` is ``int/float/bool``): - Place at function entry block begin. ``ir_value()`` creates a - fresh ``arith.constant``, so dominance is guaranteed. - - **Case B** — Block argument of the entry block (e.g., function param): - Place at function entry block begin. Block arguments dominate - everything in their block by MLIR semantics. - - **Case C** — OpResult defined in the entry block (e.g., ``a + b``): - Place immediately after the defining op via - ``InsertionPoint.after(defining_op)``. The operand is defined - right before the ref. - - **Case D** — Everything else (loop induction var, value from nested CF): - Place at current insertion point. The value normally dominates - the current point because we are using it in Python. - - **Case D-fallback** — Value from a sibling scope (e.g., stale SSA - from a previous task's ``scf.while`` that was set via Python - attribute mutation): - The SSA value no longer dominates the current insertion point. - Create a ``ub.poison``-initialised ref at the function entry - block so it is accessible from every scope. The caller will - immediately ``store`` the correct value; the poison makes any - accidental pre-store load visible as UB rather than a silent - wrong result. + **Case D-fallback** — Non-dominating sibling-scope SSA: a ``ub.poison`` ref + at function entry; the caller stores immediately (pre-store loads = UB). """ + # Some mint paths may fail to construct a stand-in cell (None) before the + # in-branch fallback re-mints; declare the union once for every branch. + mv: "MutableValue | None" entry_block = _get_function_entry_block() + # Invariant E: a place first bound inside an open region seeds at its birth + # block; entry placement stays a pure dominance device (placeholder init). + if ( + birth_block is not None + and entry_block is not None + and pyir is not None + and _block_strictly_inside(birth_block, entry_block) + ): + mv = _mint_region_born_cell(dsl_value, birth_block, entry_block) + if mv is not None: + return mv + if entry_block is not None: - # Case A: literal-backed → fresh constant at entry block begin + # Case A: literal-backed -> fresh constant at entry block begin (the mint + # site is not the binding position; the literal is the unconditional seed). if _is_literal_backed(dsl_value): with ir.InsertionPoint.at_block_begin(entry_block): mv = MutableValue(dsl_value) @@ -1328,10 +1176,8 @@ def _create_ref(dsl_value: Any) -> "MutableValue": ir_val = dsl_value.ir_value() - # Case B: block argument of entry block - # Use owner type to distinguish: Block→BlockArgument, Operation→OpResult. - # Python isinstance on ir_val is unreliable because DSL wrapper types - # (e.g. ArithValue) extend Value directly, not BlockArgument/OpResult. + # Case B: key on owner type (Block->BlockArgument, Operation->OpResult); + # isinstance on ir_val is unreliable (DSL wrappers extend Value directly). owner = ir_val.owner if isinstance(owner, ir.Block): if owner == entry_block: @@ -1350,79 +1196,82 @@ def _create_ref(dsl_value: Any) -> "MutableValue": mv.take_reference() return mv - # Determine the block where the value is defined. - # BlockArgument → owner is Block. OpResult → owner is Operation. - # NOTE: ir.Block does not support `is` identity comparison; each - # attribute access creates a new Python wrapper. Use `==` instead. - owner = ir_val.owner - if isinstance(owner, ir.Block): - defining_block = owner - else: - defining_op = getattr(owner, "operation", owner) - defining_block = defining_op.block - if entry_block is not None and pyir is not None: current_block = ir.InsertionPoint.current.block if pyir.is_value_in_ancestor_region(ir_val, current_block): - # Value dominates current IP. Check if we are in a NESTED - # region (e.g., inside scf.if) relative to the defining block. - # If so, placing the ref at the current IP traps it in that - # region — the topk bug. Use poison-init ref at function - # entry + store after the defining op instead. - # - if defining_block != current_block and _is_memref_like(dsl_value): + # A memref-backed descriptor is a recomputable value handle: its ref + # is a current-IP anchor, re-minted when a later access cannot reach it. + if _is_memref_like(dsl_value): log().info( - "[_create_ref] Case D: memref descriptor → ref at current " - "IP (no entry-block poison)" + "[_create_ref] Case D: memref descriptor -> ref at current " + "IP (recomputable value handle, no entry cell)" ) mv = MutableValue(dsl_value) mv.take_reference() return mv - if defining_block != current_block: - log().info( - "[_create_ref] Case D: poison-init ref at entry " - "block + store after defining op (nested region)" - ) - with ir.InsertionPoint.at_block_begin(entry_block): - poison_dsl = _make_poison_like(dsl_value, ir_val) - mv = MutableValue(poison_dsl) - mv.take_reference() - # Store the real value after its defining op, NOT at the - # current IP (which is inside a nested scf.if). - if isinstance(owner, ir.Block): - with ir.InsertionPoint.at_block_begin(defining_block): - mv.store(dsl_value) + # Value dominates the IP: one entry-block cell per place (never region- + # trapped), seeded by a store-after-def; init is a DEFINED zero placeholder. + log().info( + "[_create_ref] Case D: zero-init ref at entry " + "block + store after defining op" + ) + with ir.InsertionPoint.at_block_begin(entry_block): + init_ir = _make_raw_placeholder_init(ir_val.type) + if init_ir is not None: + mv = _mint_ref_with_raw_init(dsl_value, init_ir) else: - with ir.InsertionPoint.after(defining_op): - mv.store(dsl_value) - # TODO: After tracing both branches of an scf.if, verify - # that every poison-initialized ref was stored to in both - # branches. If not, raise DSLUserCodeError. This requires - # tracking which refs were created with poison init and - # which branches stored to them — complex enough to defer. - return mv - - log().info("[_create_ref] Case D: ref at current IP (same block)") - mv = MutableValue(dsl_value) - mv.take_reference() + # Non-scalar: mint the write-only placeholder cell in C++ so + # the poison init never materializes as a Python value. + mv = _mint_write_only_placeholder_cell( + dsl_value, ir_val.type, entry_block + ) + if mv is None: + # No stand-in constructible: fall back to the wrapped + # zero/poison template (pre-existing behavior). + init_dsl = _make_zero_like(dsl_value, ir_val) + mv = MutableValue(init_dsl) + mv.take_reference() + # The seed's write-fact is position + value: anchor the store at + # the defining position of the SAME raw ``mv.store`` emits (the + # value's own baked backing). ``ir_val`` can be a choke-time + # ``ir_value()`` re-serve emitted at the ambient (possibly + # region-interior) IP; anchoring there while storing the exterior + # backing leaves the entry placeholder live on every path outside + # that region, so a later region-exterior serve of the cell reads + # a value it was never given. + seed_raw = _raw_backing_ir_value(dsl_value) + if seed_raw is None or not pyir.is_value_in_ancestor_region( + seed_raw, current_block + ): + seed_raw = ir_val + seed_owner = seed_raw.owner + if isinstance(seed_owner, ir.Block): + with ir.InsertionPoint.at_block_begin(seed_owner): + mv.store(dsl_value) + else: + seed_def = getattr(seed_owner, "operation", seed_owner) + with ir.InsertionPoint.after(seed_def): + mv.store(dsl_value) return mv - # D-fallback: value does NOT dominate the current insertion point. - # This happens when a Python attribute carries a stale SSA from a - # sibling scope (e.g., task 1's scf.while value used in task 2's - # scf.while). Create a poison-initialised ref at function entry - # so it is accessible from every scope. The caller will - # immediately store the correct value. + # D-fallback (value does not dominate): poison-init ref at function entry, + # caller stores immediately; non-scalar placeholders are minted in C++. log().info( "[_create_ref] Case D-fallback: value does not dominate " "current IP → poison-init ref at entry block" ) with ir.InsertionPoint.at_block_begin(entry_block): - poison_dsl = _make_poison_like(dsl_value, ir_val) - mv = MutableValue(poison_dsl) - mv.take_reference() + mv = None + if not _pyir_type_is_scalar(ir_val.type): + mv = _mint_write_only_placeholder_cell( + dsl_value, ir_val.type, entry_block + ) + if mv is None: + poison_dsl = _make_poison_like(dsl_value, ir_val) + mv = MutableValue(poison_dsl) + mv.take_reference() return mv mv = MutableValue(dsl_value) @@ -1430,981 +1279,8635 @@ def _create_ref(dsl_value: Any) -> "MutableValue": return mv -# ---------------------------------------------------------------------- -# Slot MutableValue storage. -# -# MutableValues for dict / list / non-__dict__ / non-weakref-able owners are -# keyed by ``_make_slot_key(None, owner, slot_name)`` into the shared -# ``_slot_mvs`` table (see pyir_state.py) -- the same keying used by -# ``_slot_refs`` / ``_meta_uses``. ``_make_slot_key`` pins the owner via -# ``_pinned_owners`` so id(owner) stays stable for the trace, so no separate -# finalizer is needed. ``__dict__``-backed objects use per-object tier-1 / -# tier-2 storage instead (see ``_get_slot_mv``). -# ---------------------------------------------------------------------- +def _pyir_emit_store( + value: "ir.Value", ref: "ir.Value", *, choke: str = "", var: Any = None +) -> Any: + """The single ``pyir.store`` emission choke: refuses unfaithful stores + loudly and advances the cell's write-epoch.""" + _pointee = None + try: + _pointee = ref.type.pointee + _mismatch = value.type != _pointee + except Exception: + _mismatch = False + if _mismatch: + raise DSLRuntimeError( + "PyIR internal error: type-unfaithful pyir.store" + + (f" at {choke}" if choke else "") + + f": value type {value.type} does not equal the cell pointee " + f"type {_pointee}. Every store path must re-mint the cell at " + "the new type (function-scope redefinition) or refuse the " + "in-region type transition before emitting." + ) + # A value stored into a cell that outlives the enclosing staged region + # must not root in an in-region allocation (entry-hoisted alias): the + # escaped identity crosses the iteration boundary one generation stale. + # The judgment is TYPE-UNIFORM -- the operand-cone Allocate fact, cut at + # Read producers and block arguments, fires only for identity-carrying + # values (handles, pointers, views), never content-derived scalars. + # Register-backed memref handles go through the liveness gate (carried- + # phi admission); every other alloc-rooted value refuses -- no declared + # space/lifetime fact can prove its aliasing faithful. + _region = _pyir_enclosing_region_op_at_ip() + if ( + _region is not None + and not _ir_value_defined_inside_op(ref, _region) + and _pyir_memref_alloc_rooted_inside_region(value, _region) + ): + if not _is_memref_like(value) or not _pyir_admit_scratch_rebind( + value, ref, _region, var + ): + _pyir_raise_memref_inregion_alloc_rebind(var) + store = pyir.store(value, ref) + if _is_memref_like(value): + _pyir_note_memref_store(ref, value, getattr(store, "operation", store)) + # Bump the ref's write-epoch so cached loads are invalidated even when the + # store bypassed their ``MutableValue``. + try: + _REF_WRITE_EPOCH[ref] = _REF_WRITE_EPOCH.get(ref, 0) + 1 + except TypeError: + pass # unhashable ref stand-in -- no epoch tracking + return store -def _pyir_lookup_slot_from_value(value: "Any") -> "MutableValue | None": - """Find the ``MutableValue`` that most recently produced *value* via ``.load()``. +class _SlotId(_NamedTuple): + kind: "_Literal['scope', 'attr', 'subscript']" + owner: int # id(owner_obj), or the owning scope_id for a bare-name slot + key: Any - Cheap path: consult ``value._mutable_ref`` directly (still the - authoritative tag for staged DSL values). Returns ``None`` when - *value* is not slot-backed, so callers can route through the - plain-Python path with no auto-load. - Used by the cutlass_dsl post-loop bridge and by - ``_pyir_value_tracked_by_accessible_ref`` (the op-build dominance - check). The deliberately narrow contract (no registry-scan - fallback) keeps lookups O(1) and rules out ambiguous matches when - two slots share the same ``_load_version``. - """ - if getattr(value, "_pyir_load_version", None) is None: - return None - return getattr(value, "_mutable_ref", None) +def _purge_owner_slots(owner_id: int) -> None: + """Remove every ``_SLOT_REGISTRY`` entry owned by *owner_id* (weakref + finalizer callback on owner GC).""" + dead = [sid for sid in _SLOT_REGISTRY if sid.owner == owner_id] + for sid in dead: + _SLOT_REGISTRY.pop(sid, None) + _OWNER_KEEPALIVE.pop(owner_id, None) -def _pyir_value_tracked_by_accessible_ref(value: "Any") -> bool: - """Whether *value* is backed by a slot whose ``pyir.ref`` is reachable - from the current insertion point. +# Owner types already reported as non-weakref-able (the arm re-fires per +# write; the debug note is per-type). +_NON_WEAKREF_OWNER_TYPES_SEEN: "set[str]" = set() - Used by the op-build dominance check (``_mlir_helpers/op.py``) - to skip values whose data flow ``convert-pyir-to-scf`` will thread - through ``scf`` iter_args -- those don't need a per-call dominance - audit because the lowering pass already maintains SSA dominance. - """ - mv = _pyir_lookup_slot_from_value(value) - if mv is None: - return False - if getattr(mv, "ref", None) is None: - return False - return mv._is_ref_accessible() +def _ensure_owner_finalizer(owner: object) -> None: + """Attach a weakref to *owner* so its slot entries are purged on GC. -# --------------------------------------------------------------------------- -# D1 helpers — retroactive promotion of M values -# --------------------------------------------------------------------------- + Some built-ins (``dict``, ``list``, ``int``) reject weakrefs; their slot + entries then persist until ``_exit_function_trace`` clears the registry. + """ + owner_id = id(owner) + if owner_id in _OWNER_KEEPALIVE: + return + def _on_owner_collected(_ref: Any, oid: int = owner_id) -> None: + _purge_owner_slots(oid) -def _load_as_dsl(ref: "ir.Value", sample: Any) -> Any: - """Emit ``pyir.load(ref)`` and wrap in a DSL Numeric matching *sample*. + try: + _OWNER_KEEPALIVE[owner_id] = _weakref.ref(owner, _on_owner_collected) + except TypeError: + # Object doesn't support weakref (the normal path for exact + # dict/list); entries persist until trace exit clears the registry. + if (tname := type(owner).__name__) not in _NON_WEAKREF_OWNER_TYPES_SEEN: + _NON_WEAKREF_OWNER_TYPES_SEEN.add(tname) + log().debug( + "slot owner type %s is not weakref-able; its slot rows " + "persist until trace exit", + tname, + ) - Used by D1's promotion path to return a freshly-loaded staged value - once a slot has been promoted to ``pyir.ref``. *sample* drives the - output type: - - Python primitive (bool/int/float) -> Boolean / Int32 / Float32. - - ``_WatchedM`` wrapper -> dispatch by its ``python_value`` type. - - DSL Numeric -> rewrap via ``type(sample)(loaded_ir)``. - - Fallback: raw ``ir.Value``. +def _current_scope_id(context: str) -> int: + """Scope id of the innermost open scope frame; an empty stack is an internal + error (a scope-agnostic fallback would alias same-named locals).""" + if not _PYIR_SCOPE_STACK: + raise DSLRuntimeError( + "PyIR internal error: no function scope is open for a bare-name " + f"slot operation on '{context}'; every instrumented function body " + "opens its scope at entry, so slot operations must execute inside " + "a traced function." + ) + return _PYIR_SCOPE_STACK[-1].scope_id - Attaches a ``_mutable_ref`` (synthetic ``MutableValue``) carrying the - D1 ref so ``_pyir_auto_load_arg`` at downstream ``@dsl_user_op`` - boundaries re-emits ``pyir.load`` at the current insertion point. - Without this, a value loaded inside a for-body would carry stale SSA - after the body exits and post-CF uses would fail dominance. - """ - if pyir is None: - return None - loaded_ir = pyir.load(ref) - from .typing import Boolean, Int32, Float32, Numeric - - py_sample = sample - if isinstance(py_sample, _WatchedM): - py_sample = py_sample.python_value - - dsl_val: Any - if isinstance(py_sample, bool): - dsl_val = Boolean(loaded_ir) - elif isinstance(py_sample, int) and not isinstance(py_sample, bool): - dsl_val = Int32(loaded_ir) - elif isinstance(py_sample, float): - dsl_val = Float32(loaded_ir) - elif isinstance(py_sample, Numeric): - try: - dsl_val = type(py_sample)(loaded_ir) - except Exception: - return loaded_ir - else: - return loaded_ir - # Synthetic MutableValue so _pyir_auto_load_arg can re-load. - try: - mv = MutableValue(dsl_val) - mv._ref = ref - mv._ref_context_id = id(ir.Context.current) - _attach_mutable_ref(dsl_val, mv, "D1 _load_as_dsl") - except Exception: - pass # attach failure is non-fatal; value still works inside CF - return dsl_val +def _make_slot_id(owner: "object | None", key: Any) -> _SlotId: + """Structural slot id: bare names key on the current scope id, dict/list + owners on ('subscript', id, key), other owners on ('attr', id, name).""" + if owner is None: + return _SlotId("scope", _current_scope_id(key), key) + if isinstance(owner, (dict, list)): + _ensure_owner_finalizer(owner) + return _SlotId("subscript", id(owner), key) + _ensure_owner_finalizer(owner) + return _SlotId("attr", id(owner), key) + + +_CE_ABSENT = _Sentinel("constexpr snapshot absent") + + +def _cell_home_binding(scope_id: int, name: str) -> "tuple[int, str]": + """LangRef 3.12 section 7.13: a ``nonlocal`` name IS the nearest enclosing + scope's binding (one shared cell), so its place resolves to the cell's + HOME row. Resolution composes two declared facts -- the scope's syntactic + ``nonlocal`` set and its registered cell token; absent either fact the + scope keeps its own row (the staged-CF twin then refuses loudly).""" + if name in _PYIR_SCOPE_NONLOCAL_NAMES.get(scope_id, ()): + toks = _PYIR_SCOPE_CELL_TOKENS.get((scope_id, name)) + if toks: + # The scope's ENTRY registration is first: the freevar cell itself + # (later same-key registrations are region-frame carry cells). + home = _PYIR_CELL_HOME_BINDING.get(toks[0]) + if home is not None: + return home + return (scope_id, name) + + +def _ce_local_key(scope_id: int, name: str) -> tuple: + """F-CEPLACE local key: instance-qualified when the binding was BORN + inside a constexpr instance (open or closed -- the row stays reachable + after the instance exits, e.g. a const_expr-if-born local read after the + if); the shared key for an outside-born or pre-observation binding. + A ``nonlocal`` name first resolves to its cell's home binding.""" + scope_id, name = _cell_home_binding(scope_id, name) + owner = _CE_BINDING_OWNER.get((scope_id, name)) + if owner is not None: + return ("local", scope_id, name, owner) + return ("local", scope_id, name) + + +def _ce_note_local_assign(name: str, *, continues_binding: bool = False) -> None: + """ASSIGN choke (every ``=``/``+=`` of a bare local): maintain the + binding-BIRTH ownership fact. A live binding (outside-born ``None`` row + or an owner instance still open) keeps its birth owner -- a rebind never + changes ownership. Python semantics: a REBIND of an existing binding + (*continues_binding*, the choke read the bound value before assigning) + never creates a new binding either, so it keeps the birth owner even + after the owning constexpr instance closed. A dead-owner pure first-def + (the re-executed birth statement of a per-iteration local) births a NEW + logical binding owned by the innermost open instance.""" + from .multi_stage_manager import ( + open_constexpr_instance_serial, + open_constexpr_instance_serials, + ) + if not _PYIR_SCOPE_STACK: + return + key = _cell_home_binding(_PYIR_SCOPE_STACK[-1].scope_id, name) + rec = _CE_BINDING_OWNER.get(key, _CE_ABSENT) + if rec is not _CE_ABSENT and ( + rec is None or rec in open_constexpr_instance_serials() + ): + return # live binding: the birth fact is unchanged by a rebind + if rec is not _CE_ABSENT and continues_binding: + return # dead-owner REBIND: the one logical binding continues + _CE_BINDING_OWNER[key] = open_constexpr_instance_serial() + + +def pyir_seed_param_bindings(*names: str) -> None: + """Scope entry: DECLARE each parameter binding outside-born (``ce_owner`` + ``None``) so a later rebind inside a constexpr instance keeps the one + shared row -- the fact is seeded, never defaulted from absence. Python + semantics: a parameter binding IS a first-def at fn entry, so each param + place also carries the entry first-def-depth fact (a same-depth rebind in + this activation is straight-line, never a join; deeper rebinds refuse).""" + if not _PYIR_SCOPE_STACK: + return + scope_id = _PYIR_SCOPE_STACK[-1].scope_id + entry_depth = current_staged_cf_depth() + for name in names: + _CE_BINDING_OWNER[(scope_id, name)] = None + place = _ce_local_key(scope_id, name) + if place is not None and place not in _slot_first_def_depth_any: + _slot_first_def_depth_any[place] = entry_depth -def _ancestor_op_in_block( - op: "ir.Operation", target_block: "ir.Block | None" -) -> "ir.Operation": - """Return *op* or its nearest ancestor whose containing block is - *target_block*. - Used to hoist a per-iteration reset store out of a NESTED runtime loop - (where the earliest baked use lives) up to the block where the slot was - first-defined. Walks the op→block→owner-op chain (the same idiom as - ``_get_function_entry_block``). Falls back to *op* unchanged when - *target_block* is ``None`` or is not an ancestor (so behaviour matches - the pre-fix "insert before first use" placement). +def _make_slot_key( + target_name: "str | None", owner: Any = None, slot_name: Any = None +) -> Any: + """Canonical ledger PLACE key: locals on the (instance-qualified) + (scope_id, name) binding fact, owner slots on stable owner tokens (never + id()); ``None`` when no key is derivable.""" + if owner is not None and slot_name is not None: + return _corrected_place_for(owner, slot_name) + if target_name is not None: + return _ce_local_key(_current_scope_id(target_name), target_name) + return None - The walk terminates naturally: climbing op→block→owner-op always heads - toward the top of the IR, where the module op has no containing block - and ``cur.block`` raises. - """ - if target_block is None: - return op - cur = op - while cur is not None: - try: - blk = cur.block # ir.Operation -> containing Block - except Exception: - break # top of the IR tree: no containing block - if blk == target_block: - return cur - try: - parent = blk.owner # Block -> owning op - cur = getattr(parent, "operation", parent) - except Exception: - break - return op +def _pyir_owner_slot_is_computed(owner: Any, slot_name: Any) -> bool: + """True when *slot_name* on ``type(owner)`` is a property-family descriptor + whose value is produced by class code, not read from instance storage: no + storage place exists for the pair, so the read is anonymous (F-SHAPE). + A subscripted spelling (``coord[0]``) is judged by its base attribute -- + an element of a computed aggregate is itself computed. A tuple's storage + places are its element cells, so a named field on a tuple (a namedtuple + accessor) is always computed.""" + if isinstance(owner, (dict, list)): + return False + if isinstance(slot_name, _PlaceSeg): + base = slot_name.base + elif isinstance(slot_name, str): + base = slot_name + else: + return False + if isinstance(owner, tuple): + return True + if not isinstance(base, str): + return False # exact item-key base: no class descriptor to consult + try: + for klass in type(owner).__mro__: + desc = klass.__dict__.get(base) + if desc is None: + continue + if isinstance(desc, property): + return True + if isinstance(desc, _functools.cached_property): + # A materialized cache entry is instance storage. + inst = getattr(owner, "__dict__", None) + return not (inst is not None and base in inst) + return False + except Exception: + return False + return False -def _slot_store_for_tier1( - owner: object, *, create: bool = False -) -> "dict[Any, MutableValue] | None": - """Return the tier-1 ``__pyir_slots__`` dict on *owner*, or ``None``. - Tier-1 storage lives on ``owner.__dict__``. If *owner* does not have a - writable ``__dict__`` (built-in types, ``__slots__``-only classes, - ``dict`` / ``list``), return ``None``. When *create* is ``True`` a - fresh dict is installed via ``object.__setattr__`` so it bypasses - ``@dataclass(frozen=True)``. - """ - owner_dict = getattr(owner, "__dict__", None) - if owner_dict is None: +def _pyir_read_place(target_name: Any, owner: Any, slot_name: Any) -> Any: + """R0: the read's ledger place, derived unconditionally -- owner slots on + the token-rooted key, bare names on the (instance-qualified) local key. + Constexpr scopes gate instrumentation only, never place production: the + per-instance identity is a PRODUCED key fact (F-CEPLACE).""" + try: + return _make_slot_key( + target_name if isinstance(target_name, str) else None, owner, slot_name + ) + except Exception: return None - store = owner_dict.get(_SLOT_STORE_ATTR) - if store is None: - if not create: - return None - store = {} - try: - object.__setattr__(owner, _SLOT_STORE_ATTR, store) - except (AttributeError, TypeError): - return None - return store -def _slot_store_for_tier2( - owner: object, *, create: bool = False -) -> "dict[Any, MutableValue] | None": - """Return the tier-2 fallback entry for *owner*, or ``None``. - - Uses a module-level ``WeakKeyDictionary``. If *owner* cannot be - weakref'd (``dict``, ``list``, many built-ins), returns ``None``. - """ +def _pyir_route_is_live_place_row(mv: Any) -> bool: + """True when *mv* IS the ledger's live row for its own stamped place: the + route carries place authority (R1) -- a re-read through it is a read of + the place, not a value-identity guess.""" + place = getattr(mv, "_place", None) + if place is None: + return False try: - store = _PYIR_SLOT_FALLBACK.get(owner) - except TypeError: - # Unhashable owners -- can't key into a WeakKeyDictionary. - return None - if store is None: - if not create: - return None - store = {} - try: - _PYIR_SLOT_FALLBACK[owner] = store - except TypeError: - # Not weakref-able (e.g. dict, list). Skip tier 2 silently. - return None - return store + return _PLACE_REGISTRY.get(place) is mv + except Exception: + return False -def _slot_storage_available(owner: Any) -> bool: - """Return True if *owner* can host slot storage. +def _pyir_route_is_current(value: Any, mv: "MutableValue") -> bool: + """V-2 CURRENT: *value* still equals its cell's content, so a reload + through the route is exact. Holds for the stored template, or for a + choke-stamped product with no store since (write-epoch equal) consumed + where every open staged region predates the stamp -- a region entered + after it could re-execute a later-traced store before this read.""" + # An epoch-superseded value is a LAW-1 retained snapshot even when it is + # still the cell's bound template: an un-instrumented store (a plain + # Python callee's ``self.x += 1``) bumps the write-epoch without + # rebinding ``mv._value``, and the identity fast-path would misread that + # state as current and serve the post-store cell to a pre-store read. + _ep = _get_load_epoch(value) + if _ep is not None and mv._ref is not None and _ep != _ref_write_epoch(mv._ref): + return False + if value is mv._value: + return True + entry_tag = _get_load_region_entry(value) + return ( + entry_tag is not None + and entry_tag >= _pyir_region_entry_watermark() + and _get_load_epoch(value) == _ref_write_epoch(mv._ref) + ) - ``__dict__``-backed objects use tier-1 / tier-2 storage; ``dict`` / - ``list`` / non-weakref-able owners fall through to ``_slot_mvs``. - Every non-``None`` owner has SOME storage path available. - """ - return owner is not None +def _pyir_resolve_snapshot(value: Any, var: Any) -> Any: + """LAW-1 snapshot semantics for a routed value that is not its cell's + current binding: the value's own SSA where usable (a literal re-materializes + at its use site), else a loud refusal -- the capture is unrecoverable.""" + if _is_literal_backed(value) or _value_dominates_current_ip(value): + return value + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.SNAPSHOT_UNMATERIALIZABLE, + filename=filename, + lineno=lineno, + var=str(var), + ) -def _registry_owner(owner: Any) -> bool: - """Return True if *owner*'s slots are routed through ``_slot_mvs`` - (dict / list / non-weakrefable / non-``__dict__`` types). Regular - ``__dict__`` objects use tier-1 / tier-2 storage. - """ - if owner is None: + +def _pyir_row_binding_unobserved_write(value: Any, mv: "MutableValue") -> bool: + """V-3 COHERENT: True when the live binding at a rowed place is a DIFFERENT + cell's choke product -- proof that a write the chokes never saw rebound the + place. The row's own products stay coherent (template identity, own + route); a route-less binding names no producing cell, so it carries no + fact that could attribute a divergence. Load/store tag ints are + cell-anonymous, so cell attribution comes from the co-produced route.""" + if mv._ref is None or not _is_staged_value(value): return False - if getattr(owner, "__dict__", None) is not None: + if value is mv._value: return False - try: - import weakref as _wr + route = getattr(value, "_mutable_ref", None) + if route is None or route is mv: + return False + # The template's own route is the row-recorded content channel: a fresh + # product of it (an auto-load re-read) is the choked binding's content, + # not evidence of a foreign write. + return route is not getattr(mv._value, "_mutable_ref", None) - _wr.ref(owner) - except TypeError: - return True - return False +# --- Place-identity resolver + scope stack + reconstruction propagation --- +# Scope push/pop failures propagate loudly: a skipped push corrupts identity. -def _get_slot_mv(owner: Any, slot_name: Any) -> "MutableValue | None": - """Look up the ``MutableValue`` bound to ``(owner, slot_name)``. - ``_slot_mvs`` for dict / list owners (id-keyed); tier-1 / - tier-2 storage for ``__dict__``-backed objects. Returns ``None`` - when the slot has no recorded ``MutableValue``. +def _push_scope(kind: str, inherit: bool = False) -> "_ScopeFrame": + """Push a scope frame onto :data:`_PYIR_SCOPE_STACK` and return it. + + A ``'fn'`` scope gets a fresh monotonic ``scope_id``; a ``'region'`` scope + inherits the enclosing frame's (one place across regions; empty → fresh). """ - if owner is None: - return None - if _registry_owner(owner): - return _slot_mvs.get(_make_slot_key(None, owner, slot_name)) - store = _slot_store_for_tier1(owner) - if store is not None: - mv = store.get(slot_name) - if mv is not None: - return mv - store = _slot_store_for_tier2(owner) - if store is not None: - return store.get(slot_name) - return None + if inherit and _PYIR_SCOPE_STACK: + scope_id = _PYIR_SCOPE_STACK[-1].scope_id + else: + scope_id = next(_SCOPE_ID_COUNTER) + frame = _ScopeFrame(scope_id, kind) + _PYIR_SCOPE_STACK.append(frame) + return frame + + +def _pop_scope() -> None: + """Pop the top scope frame (no-op on an empty stack).""" + if _PYIR_SCOPE_STACK: + _PYIR_SCOPE_STACK.pop() -def _set_slot_mv(owner: Any, slot_name: Any, mv: "MutableValue") -> None: - """Record *mv* as the ``MutableValue`` for ``(owner, slot_name)``. +class _PyirScopeGuard: + """Context manager: push a scope frame on enter, pop in ``finally``. - Routes dict / list owners through ``_slot_mvs``; other owners - use tier-1 / tier-2 storage as before. See :func:`_get_slot_mv` for - the rationale. + A ``'region'`` scope inherits the enclosing ``scope_id``; a ``'fn'`` + scope mints a fresh one. Pop-in-finally means a region-builder exception + cannot leak a frame. Scope frames are load-bearing (slot keys and + region-arm paths resolve from the stack), so errors propagate. """ - if owner is None: - return - if _registry_owner(owner): - _slot_mvs[_make_slot_key(None, owner, slot_name)] = mv - return - store = _slot_store_for_tier1(owner, create=True) - if store is not None: - store[slot_name] = mv - return - store = _slot_store_for_tier2(owner, create=True) - if store is not None: - store[slot_name] = mv + __slots__ = ("_kind", "_inherit") -def _clear_slot_mv(owner: Any, slot_name: Any) -> None: - """Remove the recorded ``MutableValue`` for ``(owner, slot_name)``. + def __init__(self, kind: str = "region", inherit: bool = True) -> None: + self._kind = kind + self._inherit = inherit - Reserved for future teardown; not required by the v1 fix. Silently - does nothing when the slot has no entry. - """ - if owner is None: + def __enter__(self) -> "_PyirScopeGuard": + _push_scope(self._kind, inherit=self._inherit) + return self + + def __exit__(self, *exc: Any) -> "_Literal[False]": + _pop_scope() + return False + + +def pyir_register_scope_cells(cells_probe: Any) -> None: + """Instrumented-function entry: register the frame's closure CELLS on the + open scope. *cells_probe* is a generated zero-arg lambda referencing every + cellvar; its ``__closure__`` exposes the frame's cell objects without + evaluating any of them. A cell is the one binding a nested closure and + its defining frame share, so the (scope, name) -> cell-token rows are the + exact aliasing facts between a boundary closure read and a local write.""" + if not _PYIR_SCOPE_STACK: return - if _registry_owner(owner): - _slot_mvs.pop(_make_slot_key(None, owner, slot_name), None) + closure = getattr(cells_probe, "__closure__", None) + if not closure: return - store = _slot_store_for_tier1(owner) - if store is not None: - store.pop(slot_name, None) - store = _slot_store_for_tier2(owner) - if store is not None: - store.pop(slot_name, None) + names = getattr(getattr(cells_probe, "__code__", None), "co_freevars", ()) + scope_id = _PYIR_SCOPE_STACK[-1].scope_id + nonlocal_names = _PYIR_SCOPE_NONLOCAL_NAMES.get(scope_id, ()) + for name, cell in zip(names, closure): + tok = _owner_token(cell) + if tok is not None: + prev = _PYIR_SCOPE_CELL_TOKENS.get((scope_id, name), ()) + if tok not in prev: + _PYIR_SCOPE_CELL_TOKENS[(scope_id, name)] = prev + (tok,) + # The cell's OWNING scope registers first (its entry precedes any + # nested def's execution); a nonlocal-declared name is a freevar, + # so it never claims ownership of the shared cell. + if name not in nonlocal_names and tok not in _PYIR_CELL_HOME_BINDING: + _PYIR_CELL_HOME_BINDING[tok] = (scope_id, name) + + +def pyir_register_nonlocal_names(*names: str) -> None: + """Instrumented-function entry: DECLARE the function's ``nonlocal`` names + on the open scope (a rewrite-time syntactic fact). A bare-name choke + consults the innermost scope's set to recognize that the name binds a + rebindable closure cell of an enclosing scope, not a scope-own local.""" + if not _PYIR_SCOPE_STACK or not names: + return + scope_id = _PYIR_SCOPE_STACK[-1].scope_id + prev = _PYIR_SCOPE_NONLOCAL_NAMES.get(scope_id) + fresh = frozenset(names) + _PYIR_SCOPE_NONLOCAL_NAMES[scope_id] = fresh if prev is None else prev | fresh + + +def _pyir_boundary_cell_read_for_local(slot_key: Any) -> "tuple | None": + """The boundary closure-read record aliasing a bare-local slot key, or + ``None``. Resolution is by CELL IDENTITY: a cell token registered for the + local must equal the token recorded at the boundary read -- a same-named + local backed by a different cell never matches. Every registered cell of + the binding (entry cell plus region-frame carry cells) is consulted.""" + if not (isinstance(slot_key, tuple) and slot_key and slot_key[0] == "local"): + return None + for tok in _PYIR_SCOPE_CELL_TOKENS.get((slot_key[1], slot_key[2]), ()): + rec = _PYIR_BOUNDARY_META_CELL_READS.get(("cell", tok)) + if rec is not None: + return rec + return None -def _iter_slot_mvs_for_pyir_read( - owner: Any, -) -> "list[tuple[Any, MutableValue]]": - """Yield ``(slot_name, MutableValue)`` pairs registered for *owner*. - - Combines tier-1 (``owner.__dict__``) and tier-2 (``WeakKeyDictionary``). - Tier-1 entries take precedence: if the same slot name appears in both - tiers, only the tier-1 entry is returned (matches ``_get_slot_mv``'s - lookup order). Returns an empty list when no slots are registered or - when *owner* cannot hold slot storage. - - Used by ``pyir_read`` to refresh each registered slot on a - slot-tracked Python object before the caller (e.g. a while - before-block ``pyir_read("work_tile", work_tile)``) re-reads an - attribute via a ``@property`` getter. Without this, the property - would return the construction-time cached SSA, missing any cross- - region writes. - """ - if owner is None: - return [] - result: list[tuple[Any, "MutableValue"]] = [] - seen: set[Any] = set() - store = _slot_store_for_tier1(owner) - if store is not None: - for k, mv in store.items(): - result.append((k, mv)) - seen.add(k) - store = _slot_store_for_tier2(owner) - if store is not None: - for k, mv in store.items(): - if k not in seen: - result.append((k, mv)) - return result +def _owner_token(obj: Any) -> "int | None": + """Stable id()-independent token for slot-owner *obj*, minted on first + sight. The table is identity-keyed, so the lookup never routes through a + payload dunder (IR-silent by construction) and never raises.""" + if obj is None: + return None + tok = _OWNER_TOKENS.get(obj) + if tok is None: + tok = next(_TOKEN_COUNTER) + _OWNER_TOKENS[obj] = tok + # F-CLASS: the owner's class is recorded at token mint; carried + # uses validate the live class against it (V-7). + _PYIR_TOKEN_BORN_CLASS[tok] = type(obj) + return tok + + +def _pyir_wall_composite_str_slot(slot_name: Any) -> None: + """Bracket wall: a composite STRING slot name in attr place-key space is + impossible by construction (identifiers contain no ``[``; composite + tuple-element keys are ``_PlaceSeg``, declared at their decomposition + birth). A partial migration or regression fails LOUD here instead of + minting a silent place twin.""" + if isinstance(slot_name, str) and "[" in slot_name: + raise DSLRuntimeError( + f"[pyir internal] composite string slot name {slot_name!r} reached " + "place-key space; composite tuple-element keys must be _PlaceSeg" + ) -def _attach_mutable_ref(obj: object, mv: "MutableValue", context: str) -> None: - """Attach *mv* as ``_mutable_ref`` on *obj*. +def _place_for( + target_name: "str | None", owner: Any = None, slot_name: Any = None +) -> Any: + """Canonical place key rooted at the IMMEDIATE owner's token (locals on the + scope id); :func:`_corrected_place_for` layers stable-anchor rooting on top.""" + if owner is not None and slot_name is not None: + tok = _owner_token(owner) + if tok is None: + return None + if isinstance(owner, (dict, list)): + return ("subscript", tok, slot_name) + _pyir_wall_composite_str_slot(slot_name) + return ("attr", tok, slot_name) + if target_name is not None: + if not _PYIR_SCOPE_STACK: + return ("local", None, target_name) + return _ce_local_key(_PYIR_SCOPE_STACK[-1].scope_id, target_name) + return None - Silently logs (but does not raise) when the object does not accept - arbitrary attributes (e.g. built-in types, ``__slots__`` without - the attribute). - """ + +def _corrected_place_for(owner: Any, slot_name: Any) -> Any: + """Attr/subscript place resolver: recorded stable anchor + symbolic suffix, + else immediate-owner rooting; rooting freezes at first use.""" + if owner is None or slot_name is None: + return _place_for(None, owner, slot_name) + pref = _OWNER_PLACE_PREFIX.get(owner) + if pref is not None: + root_tok, suffix = pref + _pyir_wall_composite_str_slot(slot_name) + return ("attr", root_tok, *suffix, slot_name) + # First key for this owner: freeze its self-rooting so the key space stays + # stable for the rest of the trace. + tok = _owner_token(owner) + if tok is not None and not isinstance(owner, (dict, list)): + _OWNER_PLACE_PREFIX[owner] = (tok, ()) + return _place_for(None, owner, slot_name) + + +def _register_place_prefix( + target_name: Any, + owner: Any, + slot_name: Any, + *bound_values: Any, + alias_from: Any = None, +) -> None: + """Record the stable-anchor rooting for values bound at a dotted access + (first-wins); the *alias_from* leg lets a replacement adopt the old rooting.""" try: - object.__setattr__(obj, "_mutable_ref", mv) - except (AttributeError, TypeError): - log().info( - "could not attach _mutable_ref to %s (%s)", - type(obj).__name__, - context, + if ( + not isinstance(target_name, str) + or "." not in target_name + or owner is None + or not isinstance(slot_name, (str, _PlaceSeg)) + or _is_staged_value(owner) + ): + # A dotted attr access with a compound (non-staged) owner only; a + # staged owner must not be attribute/hash/weakref-probed. + return + segments = target_name.split(".") + # The last dotted segment must name this slot, else a synthesized + # ``field_key`` mismatch could mis-root the chain. + if len(segments) < 2 or segments[-1] != str(slot_name): + return + prefix_segments = segments[:-1] + # Exactness gate: the dotted split is exact only over identifier + # segments (an identifier contains neither '.' nor '['); a spelling + # whose subscript key carries a dot shreds into garbage segments that + # would mint a guaranteed-unresolvable rooting. Outside the exact + # domain fall through to the immediate-owner self-rooting. + if not all(s.isidentifier() for s in prefix_segments): + return + + # Resolve the owner's (root_token, owner_suffix). + existing = _OWNER_PLACE_PREFIX.get(owner) + if existing is not None: + root_tok, owner_suffix = existing + elif len(prefix_segments) == 1: + # The owner IS the outermost named binding: its token is the anchor; + # the name->token row is scoped per activation. + root_tok = _owner_token(owner) + if root_tok is None: + return + owner_suffix = () + _root_scope = _PYIR_SCOPE_STACK[-1].scope_id if _PYIR_SCOPE_STACK else -1 + _ROOT_NAME_TOKENS[(_root_scope, prefix_segments[0])] = root_tok + _OWNER_PLACE_PREFIX[owner] = (root_tok, ()) + else: + # Deeper owner not yet recorded: root at the outermost named binding if + # its token is known IN THIS SCOPE, else anchor at this owner's own token. + _root_scope = _PYIR_SCOPE_STACK[-1].scope_id if _PYIR_SCOPE_STACK else -1 + root_tok = _ROOT_NAME_TOKENS.get((_root_scope, prefix_segments[0])) + if root_tok is not None: + owner_suffix = tuple(prefix_segments[1:]) + else: + root_tok = _owner_token(owner) + if root_tok is None: + return + owner_suffix = () + + # Place-continuity rooting for a replacement rebind (resolved once): + # the old value's recorded rooting names where THIS PLACE's cells live + # (first-wins may have rooted them under an earlier spelling), so a + # structure-equal replacement adopts it -- a place fact, never an + # owner-token adoption (that channel is the rebuild protocol, V-6). + _place_rooting = None + if alias_from is not None and not is_inside_constexpr_loop(): + try: + if not _is_staged_value(alias_from) and _has_instance_storage( + alias_from + ): + _place_rooting = _OWNER_PLACE_PREFIX.get(alias_from) + except Exception: + _place_rooting = None + + leaf_suffix = owner_suffix + (slot_name,) + for bound in bound_values: + # Register only compound container objects; a staged leaf is skipped + # BEFORE any attribute/hash/weakref probe (probing it is not neutral). + if ( + bound is None + or _is_staged_value(bound) + or not _has_instance_storage(bound) + ): + continue + # First-wins over EVERY recorded rooting, the default self-rooting + # freeze included: a rooting is recorded exactly when the first + # place key is issued under it, so the record means live keys may + # already name this object's leaves. Re-rooting it mid-trace + # would land two distinct live objects' registrations on one + # place key, and the shared cell then serves one object's value + # for the other's slot (a silent wrong value). Only a bound + # object with NO recorded rooting -- a fresh replacement whose + # key space is provably unused -- may adopt the old rooting. + if _OWNER_PLACE_PREFIX.get(bound) is not None: + continue + if ( + _place_rooting is not None + and bound is not alias_from + and _pyir_same_container_structure(alias_from, bound) + ): + # Replacement rebind: the fresh object's leaves continue + # the place's existing cells (rooting recorded on the + # value it replaces). + _OWNER_PLACE_PREFIX[bound] = _place_rooting + else: + _OWNER_PLACE_PREFIX[bound] = (root_tok, leaf_suffix) + except Exception: + pass + + +def _pyir_same_container_structure(old_obj: Any, new_obj: Any) -> bool: + """True iff both objects expose the same type + instance-storage field names + (or tuple/list arity); gates owner-token propagation. Conservative: False.""" + try: + if isinstance(old_obj, (tuple, list)) and isinstance(new_obj, (tuple, list)): + return len(old_obj) == len(new_obj) + if type(old_obj) is not type(new_obj): + return False + old_d = _instance_storage_items(old_obj) + new_d = _instance_storage_items(new_obj) + if old_d is None or new_d is None: + return old_d is None and new_d is None + old_keys = {k for k in old_d if k != _SLOT_STORE_ATTR} + new_keys = {k for k in new_d if k != _SLOT_STORE_ATTR} + return old_keys == new_keys + except Exception: + return False + + +def _pyir_propagate_owner_token(old_obj: Any, new_obj: Any) -> None: + """Copy *old_obj*'s owner token onto *new_obj* (recursing into fields) so + token-rooted places re-resolve after reconstruction; structure-gated, guarded. + + Runs only inside the rebuild-protocol frame (V-6): token adoption maps a + ``__new_from_mlir_values__`` product onto its source's places; any other + fresh object is a birth and must mint its own token.""" + if not _PYIR_REBUILD_PROTOCOL_DEPTH[0]: + raise DSLRuntimeError( + "owner-token adoption outside the rebuild protocol: a token may " + "only transfer to an object produced by __new_from_mlir_values__ " + "(a constructor call births a fresh token)." ) + try: + # Iterative pair walk with a visited set: a cyclic object graph (or a + # rebuilder aliasing old sub-objects) terminates instead of recursing + # to a swallowed RecursionError with partial token propagation. + work: "list[tuple[Any, Any]]" = [(old_obj, new_obj)] + seen_pairs: "set[tuple[int, int]]" = set() + while work: + old_cur, new_cur = work.pop() + if old_cur is None or new_cur is None or old_cur is new_cur: + continue + pair = (id(old_cur), id(new_cur)) + if pair in seen_pairs: + continue + seen_pairs.add(pair) + # Register both objects as candidate holders so a rebuilt captured + # object stays discoverable by the registry-driven gathers. + _pyir_register_candidate_holder(old_cur) + _pyir_register_candidate_holder(new_cur) + _struct_ok = _pyir_same_container_structure(old_cur, new_cur) + if not _struct_ok: + continue + tok = _OWNER_TOKENS.get(old_cur) + if tok is not None: + _OWNER_TOKENS[new_cur] = tok + # Framework reconstruction is NOT a generation event: stamps + # propagate unchanged and no generation counter is bumped. + if _SUPERSEDED_GENERATIONS: + _gen_rec = _SUPERSEDED_GENERATIONS.get(id(old_cur)) + if _gen_rec is not None: + _SUPERSEDED_GENERATIONS[id(new_cur)] = _gen_rec + _pyir_keepalive_generation_obj(new_cur) + if isinstance(old_cur, (tuple, list)) and isinstance( + new_cur, (tuple, list) + ): + work.extend(zip(old_cur, new_cur)) + continue + old_d = _instance_storage_items(old_cur) + new_d = _instance_storage_items(new_cur) + if old_d is not None and new_d is not None: + for k, ov in list(old_d.items()): + if k == _SLOT_STORE_ATTR: + continue + if k in new_d: + work.append((ov, new_d[k])) + except Exception: + pass -def _fresh_wrapper(dsl_value: Any) -> Any: - """Return a fresh DSL wrapper around the same backing value as *dsl_value*. - - Used by ``pyir_assign`` at first-def to give each Python local a - distinct wrapper object. Without this, ``a = b = c = seed`` would - bind three locals to the same Python object, and the value-keyed - ``_mutable_ref`` cache used for ``ast.Name`` targets (which have no - storage owner for slot-keyed identity) would collapse onto a single - ``pyir.ref`` -- the last writer of ``seed._mutable_ref`` wins. - - Preserves the literal-vs-SSA backing so that ``_create_ref`` placement - rules (Case A for literals, Case C for SSA) behave identically to the - original wrapper. - - Returns *dsl_value* unchanged on any failure (non-staged value, - type with no single-arg constructor, etc.). Callers must continue - to treat the return value as semantically equivalent to *dsl_value*. - """ - # Literal-backed Numeric values (Int32(0), Float32(1.5), ...): build a - # fresh wrapper from the raw Python scalar so the new wrapper is also - # literal-backed. Materializing via ``ir_value()`` would emit an - # ``arith.constant`` eagerly and demote the wrapper to SSA-backed, - # which changes _create_ref placement (Case A → Case C) and shifts - # ref ordering in the IR. - if _is_literal_backed(dsl_value): - try: - return type(dsl_value)(dsl_value.value) - except Exception: - return dsl_value +def _pyir_adopt_rebuilt_owner_token(src: Any, rebuilt: Any) -> None: + """Adoption entry for the generic value-tree rebuilders: a + ``__new_from_mlir_values__`` product adopts *src*'s token inside the + rebuild-protocol frame (F-BIRTH, V-6).""" + _PYIR_REBUILD_PROTOCOL_DEPTH[0] += 1 + try: + _pyir_propagate_owner_token(src, rebuilt) + finally: + _PYIR_REBUILD_PROTOCOL_DEPTH[0] -= 1 + + +def _pyir_lookup_owner_token(obj: Any) -> "int | None": + """*obj*'s owner token WITHOUT minting, or ``None``.""" + return _OWNER_TOKENS.get(obj) + + +def _pyir_validate_owner_class(obj: Any) -> None: + """V-7: a tokenized owner's live class must equal the class recorded at its + token mint -- place keys and class facts resolved under a reclassed owner + name the wrong namespace. A never-tokenized owner is untouched.""" + if obj is None or isinstance(obj, (tuple, list, dict)): + return + tok = _pyir_lookup_owner_token(obj) + if tok is None: + return + born = _PYIR_TOKEN_BORN_CLASS.get(tok) + if born is None or type(obj) is born: + return + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.OWNER_CLASS_CHANGED, + filename=filename, + lineno=lineno, + old_class=born.__name__, + new_class=type(obj).__name__, + ) + + +def _pyir_owner_is_celled(obj: Any) -> bool: + """True when a live cell is booked under a place rooted at *obj*'s recorded + rooting: the owner carries staged state (not an effect-only object).""" + pref = _OWNER_PLACE_PREFIX.get(obj) + if pref is None: + tok = _pyir_lookup_owner_token(obj) + if tok is None: + return False + pref = (tok, ()) + root_tok, suffix = pref + n = len(suffix) + for key in _slot_refs: + if ( + isinstance(key, tuple) + and len(key) >= 2 + n + and key[0] in ("attr", "subscript") + and key[1] == root_tok + and tuple(key[2 : 2 + n]) == tuple(suffix) + ): + return True + return False + +def _pyir_emission_self_check( + op: str, + owner: Any, + slot_name: Any, + live_mv: "MutableValue | None", + place: Any, + place_mv: "MutableValue | None", +) -> None: + """Raise when a staged slot access resolves a live cell different from the + place's registered live cell of the same MLIR type (a split place).""" + if live_mv is None or place_mv is None or live_mv is place_mv: + return try: - ir_val = dsl_value.ir_value() + live_ref = live_mv._ref + place_ref = place_mv._ref + if live_ref is None or place_ref is None: + return + if not (live_mv._is_ref_accessible() and place_mv._is_ref_accessible()): + return + if live_ref.type != place_ref.type: + return except Exception: - return dsl_value + # Structural reads only: an unreadable cell cannot be classified as + # live, so it is not a violation. + return + if len(_EMISSION_SELF_CHECK_LOG) < _EMISSION_SELF_CHECK_LOG_CAP: + _EMISSION_SELF_CHECK_LOG.append( + { + "op": op, + "owner_type": type(owner).__name__, + "owner_id": id(owner), + "slot": slot_name, + "place": place, + } + ) + raise DSLRuntimeError( + "PyIR emission self-check: one-cell-per-place violation on " + f"{op} of slot {slot_name!r} (owner type {type(owner).__name__!r}, " + f"place key {place!r}): the access resolves a live cell different " + "from the place's registered cell. Two live cells of the same type " + "exist for one logical place, so stores and reads of this slot can " + "split across cells." + ) - # Mirror MutableValue._reconstruct's strategy ladder so this works - # for compound types (Array, TensorSSA) as well as scalars. - if hasattr(dsl_value, "__new_from_mlir_values__"): - try: - return dsl_value.__new_from_mlir_values__([ir_val]) - except Exception: - pass - shape = getattr(dsl_value, "_shape", None) - if shape is not None: - dtype = getattr(dsl_value, "_dtype", None) - try: - return type(dsl_value)(ir_val, shape, dtype) - except Exception: - pass +def _pyir_guard_stale_epoch(value: Any) -> None: + """Refuse an SSA-backed wrapper born in a previous, already-finalized + compilation context, before its backing ``ir.Value`` is dereferenced.""" + birth = getattr(value, "_pyir_birth_ctx", None) + if birth is None: + return + try: + current = id(ir.Context.current) + except Exception: + return + if birth != current: + from .common import DSLRuntimeError + + raise DSLRuntimeError( + "this value was produced by a previous @jit compilation and " + "cannot be reused here: its backing IR lives in an " + "already-finalized compilation context. Pass it through the " + "kernel's arguments or recompute it in this compilation." + ) + +def _pyir_witness_predicate_fold(pred: Any) -> None: + """Witness an ``if`` predicate folded on trace-time data: recording its + source places/values makes a later staged write of any of them refuse loudly.""" try: - return type(dsl_value)(ir_val) + if isinstance(pred, _WatchedM): + pred._record_structural_consumption() + return + _pyir_record_staged_literal_fold(pred) + except DSLUserCodeError: + raise # curated refusals are the loud floor, never swallowed except Exception: - return dsl_value + pass + + +def _pyir_record_staged_literal_fold(pred: Any) -> None: + """Record a literal-backed STAGED predicate fold inside an enclosing loop: + the decision is fixed at trace time, so a later staged write of any recorded + source must raise instead of silently keeping it (per-iteration hazard).""" + from .typing import Numeric + if not isinstance(pred, Numeric) or type(getattr(pred, "value", None)) not in ( + bool, + int, + float, + ): + return + # Loop multiplicity is the hazard; a loop-free fold matches Python exactly. + if not is_inside_staged_cf() or _innermost_enclosing_loop_op_at_ip() is None: + return + filename, lineno = _first_non_dsl_caller_location() + for src in getattr(pred, "_pyir_fold_srcs", ()) or (pred,): + if not isinstance(src, Numeric) or type(getattr(src, "value", None)) not in ( + bool, + int, + float, + ): + continue + if id(src) in _PYIR_STAGED_LITERAL_FOLD_WITNESSES: + continue + _PYIR_STAGED_LITERAL_FOLD_WITNESSES[id(src)] = ( + repr(src.value), + filename or "", + lineno or 0, + ) + # Pin the object so a recycled address can never alias the record. + _PYIR_STAGED_LITERAL_FOLD_KEEPALIVE.append(src) -# ========================================================================== -# MutableValue — internal bookkeeping for pyir.ref / pyir.load / pyir.store -# ========================================================================== +def _pyir_check_staged_fold_witness(value: Any, target_name: Any) -> None: + """Refuse a staged write reaching a value whose trace-time payload already + decided an if/while predicate inside an enclosing loop -- the folded branch + would keep the stale decision silently on every iteration.""" + if not _PYIR_STAGED_LITERAL_FOLD_WITNESSES: + return + rec = _PYIR_STAGED_LITERAL_FOLD_WITNESSES.get(id(value)) + if rec is None: + return + payload, read_file, read_line = rec + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.PHASE_PREDICATE_FOLDED_STALE, + filename=filename, + lineno=lineno, + var=str(target_name), + value=payload, + read_file=read_file, + read_line=read_line, + ) -class MutableValue: - """Wraps a single leaf DSL value and holds the ``pyir.ref`` handle. - The owning DSL type (Numeric) calls ``take_reference()``, - ``load()``, and ``store()`` explicitly — MutableValue is never - exposed to user code and never participates in operator dispatch. - """ +def _pyir_module_is_tracer_layer(mod: str) -> bool: + """True when *mod* is the tracer layer's own package (this module's parent, + e.g. ``cutlass.base_dsl``): its bookkeeping reads of trace-time payloads + are not consumption events. The DSL's SEMANTIC funnels (``not_``, + ``equal``, ``and_`` in the dialect layers above) consume payloads as + structure exactly like user code and are NOT exempt.""" + base_pkg = __name__.rsplit(".", 1)[0] + return mod == base_pkg or mod.startswith(base_pkg + ".") - __slots__ = ("_value", "_type", "_ref", "_ref_context_id", "_load_version") - def __init__(self, value: Any) -> None: - if isinstance(value, (bool, int, float)): - raise DSLRuntimeError( - f"Cannot create a mutable reference for Python scalar " - f"`{value}` (type: {type(value).__name__}).", - suggestion=( - "Convert to a DSL type first, e.g. " - "cutlass.Int32(...) or cutlass.Float32(...)." - ), - ) +def _pyir_record_external_payload_consumption(wrapper: Any, depth: int = 2) -> None: + """Witness a trace-time payload consumed as structure OUTSIDE the tracer + layer (a coercion dunder, comparison fold, or the payload accessor); only + the tracer's own bookkeeping reads are not consumption events. *depth* + is the consumer's frame distance: 2 for a directly-invoked dunder, 3 when + routed through one intermediate helper (the comparison funnel).""" + try: + if not _PYIR_SCOPE_STACK or not is_inside_staged_cf(): + return + # The consumer is the frame that invoked the dunder/accessor (C-level + # callees such as ``bool()``/``min()`` are frame-transparent). + mod = sys._getframe(depth).f_globals.get("__name__", "") + if _pyir_module_is_tracer_layer(mod): + return + if isinstance(wrapper, _WatchedM): + wrapper._record_structural_consumption() + else: + _pyir_record_staged_literal_fold(wrapper) + except DSLUserCodeError: + raise # curated refusals are the loud floor, never swallowed + except Exception: + pass # witnessing must never break the consumption itself - self._value = value - self._type = type(value) - self._ref = None # populated by take_reference() - self._ref_context_id: int | None = None - # Bumped on every load()/store(); used to dedup redundant - # auto-loads in ``_pyir_auto_load_arg``. A DSL value carrying - # ``_pyir_load_version`` equal to the current ``_load_version`` - # is the freshest load and need not be reloaded. - self._load_version: int = 0 - def take_reference(self) -> None: - """Create a ``pyir.ref`` for the current value. +# Trace-scoped witness: how many times a STAGED payload was consumed through +# ``__hash__`` during the active trace. Recording never refuses; the dict-key +# write choke snapshots/compares this count around hashing THE KEY, so only a +# keyed write whose own hash consumed staged identity refuses -- hashing +# anywhere else (logging, memo tables) stays untouched. +_PYIR_STAGED_HASH_WITNESS = [0] - Always creates a new ref. Caller should check - ``_is_ref_accessible()`` before calling if reuse is desired. - """ - ir_val = self._value.ir_value() - self._ref = pyir.ref(ir_val) - self._ref_context_id = id(ir.Context.current) - def _is_ref_accessible(self) -> bool: - """Return ``True`` if the existing ref is accessible from the - current insertion point (same or ancestor region).""" - if self._ref is None: - return False - if self._ref_context_id != id(ir.Context.current): - return False - current_block = ir.InsertionPoint.current.block - return pyir.is_value_in_ancestor_region(self._ref, current_block) +def _pyir_record_staged_identity_hashed() -> None: + """Witness a staged payload answering ``__hash__``: the resulting Python + hash launders runtime identity into a trace-time value, so the container + key wall must be able to see the consumption.""" + if _PYIR_SCOPE_STACK: + _PYIR_STAGED_HASH_WITNESS[0] += 1 - def _reconstruct(self, loaded_ir: Any) -> Any: - """Reconstruct a DSL value from a loaded MLIR value. - - For simple types (Int32, Float32, Boolean), ``Type(ir_value)`` - works. Complex types like ``Array`` or ``TensorSSA`` need - extra metadata (dtype, shape) that a bare ir.Value doesn't - carry. We try three strategies in order: - - 1. ``__new_from_mlir_values__`` — the extractable protocol used - by Array and similar compound types. Preserves all - internal state (dtype, shape, strides, alignment, etc.). - 2. Constructor with ``_shape``/``_dtype`` kwargs — for types - like TensorSSA that store metadata as instance attributes - and accept them as constructor kwargs. - 3. Simple ``Type(ir_value)`` — works for Numeric types (Int32, - Float32, Boolean, etc.). - """ - orig = self._value - - # Strategy 1: extractable protocol (Array, etc.) - if hasattr(orig, "__new_from_mlir_values__"): - try: - return orig.__new_from_mlir_values__([loaded_ir]) - except Exception: - pass # fall through - # Strategy 2: replay constructor with shape/dtype metadata - # (TensorSSA, Vector, etc.) - shape = getattr(orig, "_shape", None) - if shape is not None: - dtype = getattr(orig, "_dtype", None) - try: - return self._type(loaded_ir, shape, dtype) - except Exception: - pass # fall through +def _pyir_staged_hash_witness_count() -> int: + """The staged-hash witness count for snapshot/compare correlation.""" + return _PYIR_STAGED_HASH_WITNESS[0] - # Strategy 3: simple single-arg constructor - return self._type(loaded_ir) - def load(self) -> Any: - """Emit ``pyir.load`` and return a fresh DSL value. - - Tags the returned value with ``_pyir_load_version`` -- the dedup - tag ``_pyir_auto_load_arg`` checks to skip a redundant auto-load, - and the presence tag ``_pyir_lookup_slot_from_value`` requires - before returning the value's ``_mutable_ref``. - - NOTE: ``_mutable_ref`` is intentionally NOT attached here. Callers - that need ref-attachment (``pyir_assign``, ``pyir_read`` with - ``attach_ref=True``) do so explicitly via ``_attach_mutable_ref`` - after their own slot-aliasing decisions. Auto-attaching here - would break ``clone()``-style snapshots that deliberately use - ``attach_ref=False`` to prevent ref leakage across object - boundaries. - """ - assert self._ref is not None, ( - "MutableValue.load: no ref -- call take_reference() first" - ) - loaded_ir = pyir.load(self._ref) - self._load_version += 1 - loaded = self._reconstruct(loaded_ir) - try: - object.__setattr__(loaded, "_pyir_load_version", self._load_version) - except (AttributeError, TypeError): - pass # value type doesn't accept attrs — no dedup, fine. - return loaded +def _emit_constant_at_current_ip(value: Any) -> "ir.Value": + """Emit a FRESH ``arith.constant`` (bypassing the memoised helper): slots + sharing one cached SSA would clobber each other on per-slot RAUW.""" + from .._mlir.dialects import arith as _arith + from .._mlir import ir as _ir - def store(self, new_value: Any) -> None: - """Emit ``pyir.store`` to write *new_value* into the ref.""" - assert self._ref is not None, ( - "MutableValue.store: no ref -- call take_reference() first" + if isinstance(value, bool): + mlir_ty = _ir.IntegerType.get_signless(1) + return _arith.constant(mlir_ty, value) + if isinstance(value, int): + # A traced plain-int leaf materialises at the narrowest declared staged + # width that preserves its value (i32, widening to i64); anything wider + # than a signed i64 has no value-preserving width, so reject it loudly. + if -(2**31) <= value < 2**31: + return _arith.constant(_ir.IntegerType.get_signless(32), value) + if -(2**63) <= value < 2**63: + return _arith.constant(_ir.IntegerType.get_signless(64), value) + raise DSLRuntimeError( + f"PyIR: integer literal {value} read inside staged control flow " + "does not fit any declared staged integer width (i32/i64) without " + "changing its value; keep the constant within signed 64 bits or " + "restructure it as an explicit DSL integer value." ) - pyir.store(new_value.ir_value(), self._ref) - self._value = new_value - # Invalidate any previously-loaded values: subsequent uses must - # reload to observe the new value. - self._load_version += 1 - - @property - def ref(self) -> ir.Value | None: - """The raw ``pyir.ref`` SSA value (or ``None``).""" - return self._ref + if isinstance(value, float): + mlir_ty = _ir.F32Type.get() + return _arith.constant(mlir_ty, value) + # Fallback (rare): defer to the cached helper. + from .._mlir_helpers.arith import const as _arith_const - def __repr__(self) -> str: - return f"MutableValue({self._value!r})" + return _arith_const(value) -# ========================================================================== -# M→M compound auto-decomposition helpers -# ========================================================================== +def _emit_constant_for_ref(ref: "ir.Value", value: Any) -> "ir.Value": + """Emit a constant of *value* typed to match *ref*'s pointee (a store must + verify against it); falls back to the Python-default emit when unavailable.""" + from .._mlir.dialects import arith as _arith + from .._mlir import ir as _ir + try: + pointee = ref.type.pointee + except Exception: + return _emit_constant_at_current_ip(value) + try: + if _ir.IntegerType.isinstance(pointee): + width = _ir.IntegerType(pointee).width + coerced = bool(value) if width == 1 else int(value) + return _arith.constant(pointee, coerced) + # Float pointee (f16/bf16/f32/f64/...): a float attr of the exact type. + return _arith.constant(pointee, float(value)) + except Exception: + return _emit_constant_at_current_ip(value) -def _get_instance_attrs(obj: object) -> list[str]: - """Return instance attribute names from ``__dict__``. - Skips dunders. Uses ``__dict__`` (NOT ``getmembers``) to avoid - class-level properties and methods — only instance storage attributes. - """ - if not hasattr(obj, "__dict__"): - return [] - return [name for name in obj.__dict__ if not name.startswith("__")] +def _const_value_of(ir_value: "ir.Value") -> Any: + """Return the exact-typed Python value baked into an ``arith.constant`` + (implemented in C++), or ``_NO_CONST_VALUE`` when not a scalar constant.""" + try: + if pyir is None: + return _NO_CONST_VALUE + result = pyir.literal_const_of_value(ir_value) + except Exception: + return _NO_CONST_VALUE + return _NO_CONST_VALUE if result is None else result -def _is_compound_single_leaf(value: object) -> bool: - """Return True if *value* is a ref-supported value that is itself a COMPOUND - holding a staged ref-supported SEMANTIC sub-field (e.g. ``cute._Tensor``, - whose ``_iterator`` is a staged ``Pointer``). - - Such a value has a single MLIR leaf (its memref) yet is NOT a plain scalar: - threading it as a FIELD of a whole-replaced container is unsupported -- the - per-field ref it creates does not survive loop iter_args, silently dropping - the field's loop-carried value. A container holding such a field must fall - back to the clean ``CONTAINER_OBJECT_REPLACED`` rejection. (P89's top-level - whole-replace of a ``_Tensor`` itself takes the S->S path and is unaffected - by this check.) - - Scalar ref leaves (``Int32``/``Float32``/``Boolean``/``Pointer``/``Vector``) - return False: their only instance attribute is the raw MLIR ``value`` (not a - staged DSL value). The check deliberately ignores the ``_mutable_ref`` - plumbing attribute -- a scalar that has been through decomposition once - carries a ``MutableValue`` there, which must NOT make it look compound. - """ - if not (_is_staged_value(value) and _can_create_ref(value)): +def _const_values_equal(a: Any, b: Any) -> bool: + """Numeric equality that does not conflate ``bool`` with ``int`` (a baked + ``i1`` and a baked ``i32`` are distinct slot values).""" + if isinstance(a, bool) != isinstance(b, bool): return False - if not hasattr(value, "__dict__"): + try: + return a == b + except Exception: return False - for attr_name in _get_instance_attrs(value): - # Skip the scalar's own MLIR leaf (``value`` -- a bare ``ir.Value`` / - # ``ArithValue``) and the PyIR plumbing attributes. A scalar leaf - # (Int32/Float32/Boolean/Pointer/Vector) has NOTHING else; a compound - # single-leaf (``cute._Tensor``) additionally holds a SEMANTIC staged - # sub-field (``_iterator``, a ``_Pointer``) -- that is what we detect. - if ( - attr_name == "value" - or attr_name.startswith("_pyir_") - or (attr_name == "_mutable_ref") - ): - continue - sub = getattr(value, attr_name, None) - if _is_staged_value(sub): - return True - return False -def _is_leaf_decomposable(value: object) -> bool: - """Return True if *value* is a leaf that needs no further decomposition. +def _unwrap(value: Any) -> Any: + """Strip a ``_WatchedM`` wrapper (one level), or return *value* as-is.""" + if isinstance(value, _WatchedM): + return value._pyir_raw_payload + return value - A leaf is either a meta primitive (copied as-is) or a staged scalar - that ``pyir_assign`` can track via ``pyir.ref``. A ref-supported COMPOUND - single-leaf (``cute._Tensor``) is NOT a decomposable field leaf -- see - ``_is_compound_single_leaf``. + +def _pyir_structural_value_conflicts( + record: "tuple[str, str, int]", value: Any +) -> bool: + """True when writing *value* does not provably re-establish the recorded + structural consumption *record* (declared value-equality: the plain or + watched payload, or a literal-backed staged wrapper's own literal).""" + payload = _unwrap(value) + if isinstance(payload, (bool, int, float)): + return str(payload) != record[0] + literal = getattr(value, "value", None) + if isinstance(literal, (bool, int, float)): + return str(literal) != record[0] + # No provable payload (an SSA-backed staged value): the write may diverge. + return True + + +def _pyir_structural_bake_is_reseeded(slot_key: Any) -> bool: + """True when *slot_key*'s binding is re-run inside the loop body that + witnessed its structural consumption, so every iteration reaches the fold + with the value it was baked on and a later staged write cannot leave stale + structure behind:: + + while ...: + cond = False # binding: re-runs every iteration + if not cond: # fold: always reached with False -> admitted + cond = # killed before the fold is reached again + + False for everything else, which just leaves the caller's existing refusal + in place; it makes no claim about how those places are handled. + + Compares the BINDING against the CONSUMPTION, never against the write that + triggers this query: however deeply the write nests, the next iteration's + binding kills it before the fold is reached again. Testing the WRITE's + staged-CF depth instead agrees on the sketch above but rejects the SM100 + FMHA guard chains, which write from inside a staged ``if`` while their fold + stays at body level. """ - if value is None or isinstance(value, (int, float, bool, str, bytes)): - return True - if ( - _is_staged_value(value) - and _can_create_ref(value) - and not _is_compound_single_leaf(value) - ): + site = _PYIR_STRUCTURAL_META_CONSUMPTION_SITES.get(slot_key) + if site is None: + return False + sc_block, sc_loop_op = site + if sc_block is None or sc_loop_op is None: + return False + # The binding must be born inside control flow (a per-iteration re-run), not + # hoisted above the loop where its value would carry. + if not _slot_first_def_inside_cf.get(slot_key, False): + return False + def_block = _slot_first_def_block.get(slot_key) + if def_block is None: + return False + # ``is`` is unreliable: the MLIR Python bindings hand out a fresh wrapper per + # access, so two handles on the same block are distinct objects. + if def_block == sc_block: return True - return False + try: + if not _block_strictly_inside(sc_block, def_block): + return False + except Exception: + return False + # The binding must belong to the loop that witnessed the fold. A binding in + # an OUTER loop body re-runs only once per outer iteration, so an inner + # loop's fold still meets the value its own previous iteration wrote -- + # block ancestry alone would admit that. + return _block_inside_op(def_block, sc_loop_op) -def _check_tuple_decomposable( +def _pyir_merge_src_pairs(*operands: Any) -> "tuple[tuple[Any, Any], ...]": + """Union of the watched *operands*' (source place, payload) pairs, first + occurrence of a place wins (the payload its first consumption saw).""" + pairs: "list[tuple[Any, Any]]" = [] + for o in operands: + if isinstance(o, _WatchedM): + for sk, payload in o._pyir_src_pairs(): + if all(sk != have for have, _ in pairs): + pairs.append((sk, payload)) + return tuple(pairs) + + +class _WatchedM: + """Wrap a Python primitive read inside staged CF, recording its slot so a later + mutation rewrites baked constants to loads; subclasses int/float for transparency.""" + + # Declared on the base class for type-checkers; the concrete subclass __new__ + # populates each instance. + _slot_key: Any + # (source place, its payload at derivation) pairs propagated through meta + # arithmetic / comparisons (set on derived wrappers only; absent means none). + _pred_src_pairs: "tuple[tuple[Any, Any], ...]" + _cached_ir: Optional["ir.Value"] + + def __new__(cls, value: Any = 0, slot_key: Any = None) -> "_WatchedM": + # Factory dispatch: _WatchedM(...) forwards to the right int/float-backed + # subclass; subclass __new__ paths (cls is not _WatchedM) bypass this. + if cls is _WatchedM: + # Bool first (``isinstance(True, int)`` is True): route via ``_WatchedBool`` so + # ``python_value`` is True/False and ``arith.const`` emits ``i1``, not ``i32``. + if isinstance(value, bool): + return _WatchedBool.__new__(_WatchedBool, value, slot_key) + if isinstance(value, int): + return _WatchedInt.__new__(_WatchedInt, value, slot_key) + if isinstance(value, float): + return _WatchedFloat.__new__(_WatchedFloat, value, slot_key) + raise TypeError( + f"_WatchedM cannot wrap value of type {type(value).__name__}; " + "only bool/int/float are supported." + ) + # Subclass __new__ already constructed the instance; nothing to do. + return super().__new__(cls) + + def _pyir_record_birth_facts(self, slot_key: Any) -> None: + """Birth fact: the slot's cell (and write epoch) as of creation -- the + read position this wrapper snapshots (C5 one-semantics anchor).""" + ref = _slot_refs.get(slot_key) if slot_key is not None else None + self._pyir_birth_ref = ref + self._pyir_birth_epoch = _ref_write_epoch(ref) if ref is not None else None + + @property + def python_value(self) -> Any: + """Declared payload accessor: a read from outside the DSL package is a + recorded structural consumption; subclasses supply the raw payload.""" + _pyir_record_external_payload_consumption(self) + return self._pyir_raw_payload + + @property + def _pyir_raw_payload(self) -> Any: + """The bare payload for the tracer's own bookkeeping (never recorded); + subclasses cast to their int/float base.""" + return self + + def ir_value( + self, + *, + loc: "ir.Location | None" = None, + ip: "ir.InsertionPoint | None" = None, + ) -> "ir.Value": + """Emit (or reuse, if it still dominates the current IP) the leaf constant + and record it in ``_meta_uses[slot]``; ``loc``/``ip`` accepted and ignored.""" + # An arm-locally mutated slot baked outside its arm would freeze one + # traced path's value for both runtime paths -- refuse loudly. + if self._slot_key is not None: + _pyir_check_arm_local_escape(self._slot_key) + cached = self._cached_ir + if cached is not None: + # Follow the promotion rewrite: new consumers must use the RAUW'd + # position-correct load, not the dead constant. + _replacement = _META_CONST_REPLACEMENTS.get(cached) + if _replacement is not None: + cached = _replacement + self._cached_ir = cached + if cached is not None and _cached_ir_value_dominates_current_ip(cached): + return cached + # An already-promoted slot materialises as a load of its cell -- but + # only while the cell still holds this wrapper's creation-time value + # (same ref, same write epoch): one declared SNAPSHOT semantics. + if self._slot_key is not None: + _ref = _slot_refs.get(self._slot_key) + if isinstance(_ref, ir.Value) and _cached_ir_value_dominates_current_ip( + _ref + ): + if _ref is getattr(self, "_pyir_birth_ref", None) and _ref_write_epoch( + _ref + ) == getattr(self, "_pyir_birth_epoch", None): + loaded = pyir.load(_ref) + self._cached_ir = loaded + return loaded + # The cell moved past this wrapper's snapshot. A flat region + # runs once, so the payload constant IS the creation-time + # value; inside a loop the carry engine's store-back walk + # re-presents slots through this arm, so the cell load stays + # (the user-facing loop-snapshot residue is a recorded + # deviation, not a refusal). + if _innermost_enclosing_loop_op_at_ip() is None: + const = _emit_constant_at_current_ip(self._pyir_raw_payload) + self._cached_ir = const + # F-SPEC: a baked constant of a rooted place is a + # specialization fact of this trace. + _pyir_spec_record_read(self._slot_key, self._pyir_raw_payload) + return const + loaded = pyir.load(_ref) + self._cached_ir = loaded + return loaded + const = _emit_constant_at_current_ip(self._pyir_raw_payload) + self._cached_ir = const + if self._slot_key is not None: + _meta_uses.setdefault(self._slot_key, []).append(const) + # F-SPEC: a baked constant of a rooted place is a specialization fact. + _pyir_spec_record_read(self._slot_key, self._pyir_raw_payload) + return const + + def _record_structural_consumption(self) -> None: + """Witness one structural (``__index__``) consumption inside staged CF: such + a bake has no retargetable SSA constant, so a later promotion refuses loudly.""" + try: + pairs = self._pyir_src_pairs() + if not pairs or not is_inside_staged_cf(): + return + # Loop multiplicity is the hazard (per-trace bake vs per-iteration + # re-evaluation); outside any enclosing loop record nothing. + _sc_loop_op = _innermost_enclosing_loop_op_at_ip() + if _sc_loop_op is None: + return + try: + _sc_block = ir.InsertionPoint.current.block + except Exception: + _sc_block = None + filename, lineno = _first_non_dsl_caller_location() + for sk, payload in pairs: + # An arm-locally mutated slot consumed outside its arm is + # path-dependent -- refuse before witnessing the bake. + _pyir_check_arm_local_escape(sk) + # F-SPEC: a structural bake of a rooted place is a + # specialization fact of this trace (the PLACE's own payload). + _pyir_spec_record_read(sk, payload) + if sk in _PYIR_STRUCTURAL_META_CONSUMPTIONS: + continue + _PYIR_STRUCTURAL_META_CONSUMPTIONS[sk] = ( + str(payload), + filename or "", + lineno or 0, + ) + _PYIR_STRUCTURAL_META_CONSUMPTION_SITES[sk] = (_sc_block, _sc_loop_op) + except DSLUserCodeError: + raise # curated refusals are the loud floor, never swallowed + except Exception: + pass # witnessing must never break the consumption itself + + def _pyir_src_pairs(self) -> "tuple[tuple[Any, Any], ...]": + """(place key, that place's payload) pairs this value's Python value + derives from: its own slot at its own payload, plus propagated sources + each at the payload it held when the derivation consumed it.""" + pairs: "list[tuple[Any, Any]]" = [] + if self._slot_key is not None: + pairs.append((self._slot_key, self._pyir_raw_payload)) + for sk, payload in getattr(self, "_pred_src_pairs", ()): # derived + if sk is not None and all(sk != have for have, _ in pairs): + pairs.append((sk, payload)) + return tuple(pairs) + + # Coercion record arms: the closed dunder set through which non-DSL code + # consumes the payload as structure (each records, then coerces). + def __bool__(self) -> bool: + _pyir_record_external_payload_consumption(self) + return bool(self._pyir_raw_payload) + + def __int__(self) -> int: + _pyir_record_external_payload_consumption(self) + return int(self._pyir_raw_payload) + + def __float__(self) -> float: + _pyir_record_external_payload_consumption(self) + return float(self._pyir_raw_payload) + + def __hash__(self) -> int: + _pyir_record_external_payload_consumption(self) + return hash(self._pyir_raw_payload) + + def __str__(self) -> str: + _pyir_record_external_payload_consumption(self) + return str(self._pyir_raw_payload) + + def __format__(self, format_spec: str) -> str: + _pyir_record_external_payload_consumption(self) + return format(self._pyir_raw_payload, format_spec) + + def __reduce__(self) -> Any: + # Snapshot semantics (same contract as the watched containers): + # serialization de-watches to the PLAIN payload -- a recorded + # structural consumption, never a pickled tracer wrapper. + _pyir_record_external_payload_consumption(self) + payload = self._pyir_raw_payload + return (type(payload), (payload,)) + + def _pyir_cmp_watched(self, other: Any, op: str) -> Any: + """Comparison on the Python values, returning a place-attributed + ``_WatchedBool``; equality routes here like the orderings (structural + refusals stay value-conflict-gated downstream).""" + other_py = _unwrap(other) + if type(other_py) not in (bool, int, float): + return NotImplemented + import operator as _op + + result = getattr(_op, op)(self._pyir_raw_payload, other_py) + # A comparison folds the payloads (no eager-IR form): witness the + # structural bake through the consumption funnel (depth 3 = the + # frame that invoked the comparison dunder), so a later carry or + # promotion of a source place with a changed value refuses instead + # of keeping the folded truth. Only the tracer layer's own + # bookkeeping comparisons are exempt; the DSL's semantic funnels + # (``equal``, ``not_``, ``and_``, ...) consume the payload as + # structure exactly like user code does. + _pyir_record_external_payload_consumption(self, depth=3) + if isinstance(other, _WatchedM): + _pyir_record_external_payload_consumption(other, depth=3) + wrapped = _WatchedM(bool(result), None) + wrapped._pred_src_pairs = _pyir_merge_src_pairs(self, other) + return wrapped + + def __lt__(self, other: Any) -> Any: + return self._pyir_cmp_watched(other, "lt") + + def __le__(self, other: Any) -> Any: + return self._pyir_cmp_watched(other, "le") + + def __gt__(self, other: Any) -> Any: + return self._pyir_cmp_watched(other, "gt") + + def __ge__(self, other: Any) -> Any: + return self._pyir_cmp_watched(other, "ge") + + def __eq__(self, other: Any) -> Any: + return self._pyir_cmp_watched(other, "eq") + + def __ne__(self, other: Any) -> Any: + return self._pyir_cmp_watched(other, "ne") + + # ``__repr__`` and the remaining int/float dunders are inherited from the + # int/float base class on the concrete subclass. + + # ----- arithmetic: emit IR eagerly, return derived wrappers ----------- + + def _binop_ir(self, other: Any, op_name: str) -> Any: + """Emit an arith op and return a derived ``_WatchedM`` wrapper.""" + try: + from .._mlir.dialects import arith as _arith + except ImportError: + return NotImplemented + + rhs_py = _unwrap(other) + if not isinstance(rhs_py, (bool, int, float)): + return NotImplemented + + py_ops = { + "add": lambda a, b: a + b, + "sub": lambda a, b: a - b, + "mul": lambda a, b: a * b, + "truediv": lambda a, b: a / b, + "floordiv": lambda a, b: a // b, + "mod": lambda a, b: a % b, + "and": lambda a, b: a & b, + "or": lambda a, b: a | b, + "xor": lambda a, b: a ^ b, + "lshift": lambda a, b: a << b, + "rshift": lambda a, b: a >> b, + } + if op_name not in py_ops: + return NotImplemented + try: + result_py = py_ops[op_name](self._pyir_raw_payload, rhs_py) + except ZeroDivisionError: + raise # Python truth: the operation itself raises at trace time. + except Exception: + return NotImplemented + + try: + from .._mlir_helpers.arith import const as _arith_const + + lhs_ir = self.ir_value() + # SSA == payload (Python semantics): i1 operands widen to i32 + # before arithmetic (Python bool arithmetic is int arithmetic); + # bitwise ops on two bools stay i1 (Python returns bool). + _widen = ( + isinstance(lhs_ir.type, ir.IntegerType) + and lhs_ir.type.width == 1 + and not (op_name in ("and", "or", "xor") and isinstance(rhs_py, bool)) + ) + if _widen: + lhs_ir = _arith.extui(ir.IntegerType.get_signless(32), lhs_ir) + if op_name == "truediv" and isinstance(lhs_ir.type, ir.IntegerType): + # Python ``/`` always yields a float: run in the f32 domain. + lhs_ir = _arith.sitofp(ir.F32Type.get(), lhs_ir) + if isinstance(other, _WatchedM): + rhs_ir = other.ir_value() + if ( + isinstance(rhs_ir.type, ir.IntegerType) + and rhs_ir.type.width == 1 + and isinstance(lhs_ir.type, ir.IntegerType) + and lhs_ir.type.width != 1 + ): + rhs_ir = _arith.extui(lhs_ir.type, rhs_ir) + else: + rhs_ir = _arith_const(rhs_py, lhs_ir.type) + is_float = isinstance(lhs_ir.type, ir.FloatType) + arith_fn_table = { + ("add", False): _arith.addi, + ("add", True): _arith.addf, + ("sub", False): _arith.subi, + ("sub", True): _arith.subf, + ("mul", False): _arith.muli, + ("mul", True): _arith.mulf, + ("floordiv", False): _arith.floordivsi, + ("and", False): _arith.andi, + ("or", False): _arith.ori, + ("xor", False): _arith.xori, + ("lshift", False): _arith.shli, + ("rshift", False): _arith.shrsi, + } + if op_name == "truediv": + # The lhs is already float (widened above for int payloads); + # widen an integer rhs the same way (unsigned for i1). + if isinstance(rhs_ir.type, ir.IntegerType): + if rhs_ir.type.width == 1: + rhs_ir = _arith.extui(ir.IntegerType.get_signless(32), rhs_ir) + rhs_ir = _arith.sitofp(lhs_ir.type, rhs_ir) + result_ir = _arith.divf(lhs_ir, rhs_ir) + elif op_name == "floordiv" and is_float: + # Python float floordiv FLOORS; arith.divf alone does not. + from .._mlir.dialects import math as _math + + result_ir = _math.floor(_arith.divf(lhs_ir, rhs_ir)) + elif op_name == "mod" and not is_float: + # Python % takes the DIVISOR's sign; arith.remsi the dividend's. + # r = remsi(a,b); fix = (r != 0) && (sign(r) != sign(b)) -> r+b. + _r = _arith.remsi(lhs_ir, rhs_ir) + _zero = _arith_const(0, _r.type) + _r_neg = _arith.cmpi(_arith.CmpIPredicate.slt, _r, _zero) + _b_neg = _arith.cmpi(_arith.CmpIPredicate.slt, rhs_ir, _zero) + _sign_diff = _arith.xori(_r_neg, _b_neg) + _r_nonzero = _arith.cmpi(_arith.CmpIPredicate.ne, _r, _zero) + _fix = _arith.andi(_sign_diff, _r_nonzero) + result_ir = _arith.select(_fix, _arith.addi(_r, rhs_ir), _r) + elif op_name == "mod" and is_float: + # Python fmod takes the divisor's sign: r = remf; fix as above. + _r = _arith.remf(lhs_ir, rhs_ir) + _zero = _arith_const(0.0, _r.type) + _r_neg = _arith.cmpf(_arith.CmpFPredicate.OLT, _r, _zero) + _b_neg = _arith.cmpf(_arith.CmpFPredicate.OLT, rhs_ir, _zero) + _sign_diff = _arith.xori(_r_neg, _b_neg) + _r_nonzero = _arith.cmpf(_arith.CmpFPredicate.ONE, _r, _zero) + _fix = _arith.andi(_sign_diff, _r_nonzero) + result_ir = _arith.select(_fix, _arith.addf(_r, rhs_ir), _r) + else: + fn = arith_fn_table.get((op_name, is_float)) + if fn is None: + return NotImplemented + result_ir = fn(lhs_ir, rhs_ir) + except Exception: + return NotImplemented + + derived = _WatchedM(result_py, slot_key=None) + derived._cached_ir = result_ir + # Source-place propagation: a structural consumption of the derived + # value bakes the operands' places at their consumed payloads. + pairs = _pyir_merge_src_pairs(self, other) + if pairs: + derived._pred_src_pairs = pairs + return derived + + def __add__(self, other: Any) -> Any: + return self._binop_ir(other, "add") + + def __radd__(self, other: Any) -> Any: + if not isinstance(_unwrap(other), (bool, int, float)): + return NotImplemented + return _WatchedM(_unwrap(other)).__add__(self) + + def __sub__(self, other: Any) -> Any: + return self._binop_ir(other, "sub") + + def __rsub__(self, other: Any) -> Any: + if not isinstance(_unwrap(other), (bool, int, float)): + return NotImplemented + return _WatchedM(_unwrap(other)).__sub__(self) + + def __mul__(self, other: Any) -> Any: + return self._binop_ir(other, "mul") + + def __rmul__(self, other: Any) -> Any: + return self.__mul__(other) + + def __truediv__(self, other: Any) -> Any: + return self._binop_ir(other, "truediv") + + def __rtruediv__(self, other: Any) -> Any: + if not isinstance(_unwrap(other), (bool, int, float)): + return NotImplemented + return _WatchedM(_unwrap(other)).__truediv__(self) + + def __floordiv__(self, other: Any) -> Any: + return self._binop_ir(other, "floordiv") + + def __mod__(self, other: Any) -> Any: + return self._binop_ir(other, "mod") + + def __and__(self, other: Any) -> Any: + return self._binop_ir(other, "and") + + def __or__(self, other: Any) -> Any: + return self._binop_ir(other, "or") + + def __xor__(self, other: Any) -> Any: + return self._binop_ir(other, "xor") + + def __lshift__(self, other: Any) -> Any: + return self._binop_ir(other, "lshift") + + def __rshift__(self, other: Any) -> Any: + return self._binop_ir(other, "rshift") + + def __divmod__(self, other: Any) -> Any: + # LangRef 3.3.8: divmod == (floordiv, mod). Route both halves through + # the IR-emitting arm so the pair follows a later promotion of the place. + q = self._binop_ir(other, "floordiv") + r = self._binop_ir(other, "mod") if q is not NotImplemented else q + if q is NotImplemented or r is NotImplemented: + rhs = _unwrap(other) + if not isinstance(rhs, (bool, int, float)): + return NotImplemented + # Payload-domain fallback (no IR context / zero divisor): witness + # the structural consumption, then let Python decide. + self._record_structural_consumption() + if isinstance(other, _WatchedM): + other._record_structural_consumption() + return divmod(self._pyir_raw_payload, rhs) + return (q, r) + + def __rdivmod__(self, other: Any) -> Any: + lhs = _unwrap(other) + if not isinstance(lhs, (bool, int, float)): + return NotImplemented + return _WatchedM(lhs).__divmod__(self) + + def __pow__(self, other: Any, mod: Any = None) -> Any: + """Witness arm (LangRef 3.3.8): pow has no eager-IR form, so compute + the Python truth, record the structural bake, propagate source places.""" + rhs = _unwrap(other) + if not isinstance(rhs, (bool, int, float)): + return NotImplemented + if mod is not None and not isinstance(_unwrap(mod), (bool, int, float)): + return NotImplemented + for o in (self, other, mod): + if isinstance(o, _WatchedM): + o._record_structural_consumption() + if mod is None: + result = pow(self._pyir_raw_payload, rhs) + else: + result = pow(self._pyir_raw_payload, rhs, _unwrap(mod)) + if not isinstance(result, (bool, int, float)): + return result # complex etc.: the payload leaves the numeric domain + wrapped = _WatchedM(result, None) + pairs = _pyir_merge_src_pairs(self, other, mod) + if pairs: + wrapped._pred_src_pairs = pairs + return wrapped + + def __rpow__(self, other: Any, mod: Any = None) -> Any: + lhs = _unwrap(other) + if not isinstance(lhs, (bool, int, float)): + return NotImplemented + return _WatchedM(lhs).__pow__(self, mod) + + def __complex__(self) -> complex: + _pyir_record_external_payload_consumption(self) + return complex(self._pyir_raw_payload) + + # ----- unary operators: emit IR eagerly, return derived wrappers ------- + + def _unary_ir(self, op_name: str) -> Any: + """Python-truth unary fold plus the eager arith form; a failed emission + (no MLIR context) keeps the payload-only derived wrapper.""" + _py_unary = { + "neg": lambda a: -a, + "pos": lambda a: +a, + "invert": lambda a: ~a, + "abs": abs, + } + # Python truth decides validity (e.g. ``~`` on a float raises here). + result_py = _py_unary[op_name](self._pyir_raw_payload) + result_ir = None + try: + from .._mlir.dialects import arith as _arith + from .._mlir_helpers.arith import const as _arith_const + + v = self.ir_value() + if isinstance(v.type, ir.IntegerType) and v.type.width == 1: + # Python unary arithmetic on bool runs in the int domain. + v = _arith.extui(ir.IntegerType.get_signless(32), v) + is_float = isinstance(v.type, ir.FloatType) + if op_name == "pos": + result_ir = v + elif op_name == "neg": + zero = _arith_const(0.0 if is_float else 0, v.type) + result_ir = _arith.subf(zero, v) if is_float else _arith.subi(zero, v) + elif op_name == "invert": + result_ir = _arith.xori(v, _arith_const(-1, v.type)) + elif is_float: # abs + from .._mlir.dialects import math as _math + + result_ir = _math.absf(v) + else: # abs, integer domain + zero = _arith_const(0, v.type) + neg = _arith.subi(zero, v) + isneg = _arith.cmpi(_arith.CmpIPredicate.slt, v, zero) + result_ir = _arith.select(isneg, neg, v) + except Exception: + result_ir = None # payload-only wrapper: value-exact, no eager IR + derived = _WatchedM(result_py, None) + if result_ir is not None: + derived._cached_ir = result_ir + pairs = _pyir_merge_src_pairs(self) + if pairs: + derived._pred_src_pairs = pairs + return derived + + def __neg__(self) -> Any: + return self._unary_ir("neg") + + def __pos__(self) -> Any: + return self._unary_ir("pos") + + def __abs__(self) -> Any: + return self._unary_ir("abs") + + def __invert__(self) -> Any: + return self._unary_ir("invert") + + +class _WatchedInt(_WatchedM, int): + """Concrete D1 wrapper backed by ``int`` (no ``__slots__``: CPython forbids + them on int subclasses; the ``__dict__`` cost buys isinstance transparency).""" + + def __new__(cls, value: Any, slot_key: Any = None) -> "_WatchedInt": + inst = int.__new__(cls, value) + inst._slot_key = slot_key + inst._cached_ir = None + inst._pyir_record_birth_facts(slot_key) + return inst + + @property + def _pyir_raw_payload(self) -> int: + # Base-slot extraction: int(self) would re-enter the recorded __int__. + return int.__int__(self) + + def __index__(self) -> int: + # Witness the structural bake so a later promotion of the place + # refuses instead of leaving it silently stale. + self._record_structural_consumption() + return int.__int__(self) + + +class _WatchedBool(_WatchedM, int): + """Concrete D1 wrapper for ``bool`` (final in CPython, so backed by ``int`` + with ``python_value`` returning a proper ``bool`` for ``i1`` emission).""" + + def __new__(cls, value: Any, slot_key: Any = None) -> "_WatchedBool": + inst = int.__new__(cls, int(value)) + inst._slot_key = slot_key + inst._cached_ir = None + inst._pyir_record_birth_facts(slot_key) + return inst + + @property + def _pyir_raw_payload(self) -> bool: + # Base-slot extraction: int(self) would re-enter the recorded __int__. + return bool(int.__int__(self)) + + def __index__(self) -> int: + # See ``_WatchedInt.__index__`` -- same structural-bake witness. + self._record_structural_consumption() + return int.__int__(self) + + +class _WatchedFloat(_WatchedM, float): + """Concrete D1 wrapper backed by ``float``. Same ``__slots__`` rule + as :class:`_WatchedInt`.""" + + def __new__(cls, value: Any, slot_key: Any = None) -> "_WatchedFloat": + inst = float.__new__(cls, value) + inst._slot_key = slot_key + inst._cached_ir = None + inst._pyir_record_birth_facts(slot_key) + return inst + + @property + def _pyir_raw_payload(self) -> float: + # Base-slot extraction: float(self) would re-enter the recorded __float__. + return float.__float__(self) + + +# --- Watched dict: the OBJECT-level choke for dict-entry places --- + + +# Keys CREATED through the watched write choke inside dynamic staged CF +# (first-def subscript insertion stays supported: the cell is minted and item +# reads refuse staleness). The key SET, however, changed on the one traced +# pass only -- a later whole-key-set consumption (membership, lookup miss) +# would bake that pass's truth regardless of the runtime branch, so it +# refuses through this record. id -> (container, created keys); trace-scoped. +_PYIR_DICT_CF_CREATED_KEYS: "dict[int, tuple[Any, set]]" = {} + + +class _WatchedDict(dict): + """Identity-adopted tracked ``dict``: entry accesses from ANY Python code route + through one owner-keyed choke; inert without an open PyIR trace scope.""" + + # Slot attribute types (assigned post-construction at the adoption / + # write chokes, so declared here for the type checker). + _pyir_label: str + _pyir_staged_writes: bool + # The instance whose attribute storage this mapping IS (weakref, or the + # instance itself when non-weakrefable); set at the __dict__ adoption. + _pyir_instance_of: Any + + __slots__ = ("_pyir_label", "_pyir_staged_writes", "_pyir_instance_of") + + def _pyir_engaged(self) -> bool: + """Whether the object-side chokes fire for this access (see class doc).""" + return ( + not _WATCHED_CONTAINER_BYPASS[0] + and bool(_PYIR_SCOPE_STACK) + and _WATCHED_DICT_READ_HOOK[0] is not None + ) + + def __getitem__(self, key: Any) -> Any: + try: + val = dict.__getitem__(self, key) + except KeyError: + # A missed lookup consumed the key SET (the absence is the fact + # the trace acts on): record the contents snapshot, then raise. + self._pyir_record_miss() + raise + if not self._pyir_engaged(): + return val + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + return _WATCHED_DICT_READ_HOOK[0](self, key, val) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + + def _pyir_refuse_baked_keyset_read(self) -> None: + """Refuse a whole-key-set consumption after a key was CREATED inside + dynamic staged CF: the consumed key set reflects the one traced pass, + not the runtime branch that decides the insertion.""" + entry = _PYIR_DICT_CF_CREATED_KEYS.get(id(self)) + if entry is None: + return + raise DSLUserCodeError( + DiagId.CONTAINER_DICT_KEY_SET_BAKED_READ, + var=getattr(self, "_pyir_label", None) or "dict", + detail=", ".join(sorted(repr(k) for k in entry[1])), + ) + + def _pyir_record_miss(self) -> None: + # A key-set consumption with no item choke (F-SPEC, same funnel as + # ``__contains__``). + if self._pyir_engaged(): + self._pyir_refuse_baked_keyset_read() + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + _pyir_spec_record_membership(self) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + + def get(self, key: Any, default: Any = None) -> Any: + if not dict.__contains__(self, key): + self._pyir_record_miss() + if self._pyir_engaged() and _WATCHED_DICT_GET_MISS_HOOK[0] is not None: + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + _WATCHED_DICT_GET_MISS_HOOK[0](self, key) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + return default + return self.__getitem__(key) + + def __contains__(self, key: Any) -> bool: + # Membership consumes the WHOLE key set through the C-level probe + # (no item choke fires): record a contents-snapshot row (F-SPEC). + if self._pyir_engaged(): + self._pyir_refuse_baked_keyset_read() + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + _pyir_spec_record_membership(self) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + return dict.__contains__(self, key) + + def __setitem__(self, key: Any, value: Any) -> None: + if not self._pyir_engaged(): + dict.__setitem__(self, key, value) + return + # Diagnostics-only source attribution via the ONE sanctioned frame + # helper (never influences emission or slot identity). + filename, lineno = _first_non_dsl_caller_location() + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + value = _WATCHED_DICT_WRITE_HOOK[0]( + self, key, value, filename or "", lineno or 0 + ) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + dict.__setitem__(self, key, value) + + # Structural mutators (key-set changes) are trace-time-only: the object-side + # hook refuses them inside dynamic staged CF and realizes them everywhere else. + + def _pyir_structural(self, op_name: str, detail: str) -> None: + if self._pyir_engaged() and _WATCHED_DICT_MUTATOR_HOOK[0] is not None: + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + _WATCHED_DICT_MUTATOR_HOOK[0](self, op_name, detail) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + + def __iter__(self) -> Any: + # Bare iteration walks the KEY set, which is trace-time stable inside + # dynamic staged CF (structural mutation refuses at the mutator + # choke), so the raw walk is exact -- per-slot reads re-enter the + # item choke themselves. The object-side hook refuses the one + # unstable case (a key CREATED inside the region) and raw-delegates + # everywhere else. The ``.keys()/.values()/.items()`` views route at + # the call boundary (user calls only), so DSL-internal walks over + # adopted ``__dict__`` mappings stay unobserved. + if self._pyir_engaged() and _WATCHED_DICT_ITER_HOOK[0] is not None: + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + _WATCHED_DICT_ITER_HOOK[0](self, "iter()") + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + return dict.__iter__(self) + + def __delitem__(self, key: Any) -> None: + self._pyir_structural("__delitem__", f"del [{key!r}]") + dict.__delitem__(self, key) + + def _pyir_read_out(self, key: Any, val: Any) -> Any: + """Route a departing value (``pop``/``popitem``) through the entry's read + lifecycle so consumers observe the place cell, not region-trapped SSA.""" + if not self._pyir_engaged(): + return val + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + return _WATCHED_DICT_READ_HOOK[0](self, key, val) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + + def pop(self, *args: Any) -> Any: + self._pyir_structural("pop", "pop()") + if args and dict.__contains__(self, args[0]): + val = self._pyir_read_out(args[0], dict.__getitem__(self, args[0])) + dict.__delitem__(self, args[0]) + return val + return dict.pop(self, *args) + + def popitem(self) -> Any: + self._pyir_structural("popitem", "popitem()") + key, val = dict.popitem(self) + return key, self._pyir_read_out(key, val) + + def clear(self) -> None: + self._pyir_structural("clear", "clear()") + dict.clear(self) + + def setdefault(self, key: Any, default: Any = None) -> Any: + if dict.__contains__(self, key): + return self.__getitem__(key) + self._pyir_structural("setdefault", f"setdefault({key!r})") + dict.__setitem__(self, key, default) + return default + + def update(self, *args: Any, **kwargs: Any) -> None: + other: dict = {} + if args: + other.update(args[0]) + other.update(kwargs) + new_keys = [k for k in other if not dict.__contains__(self, k)] + if new_keys: + self._pyir_structural("update", f"update() inserting {new_keys!r}") + for k in new_keys: + dict.__setitem__(self, k, other.pop(k)) + # Existing keys route through the per-entry write choke, emitting the + # same stores subscript assignment would. + for k, v in other.items(): + self[k] = v + + # ``|`` keeps dict.__or__ (a PLAIN dict copy by design), so the augmented + # form's in-place signature intentionally diverges from it. + def __ior__(self, other: Any) -> "_WatchedDict": # type: ignore[misc] + self.update(other) + return self + + # --- snapshot semantics: copies are PLAIN dicts -------------------------- + + def __copy__(self) -> dict: + return dict(self) + + def __deepcopy__(self, memo: Any) -> dict: + import copy as _copy + + out: dict = {} + memo[id(self)] = out + for k, v in dict.items(self): + out[_copy.deepcopy(k, memo)] = _copy.deepcopy(v, memo) + return out + + def __reduce__(self) -> Any: + # Serialization de-watches to a PLAIN dict; while engaged, each entry + # is read THROUGH the choke so the whole-container bake records rows + # (a staged entry then reaches its natively unpicklable MLIR payload). + if self._pyir_engaged(): + return (dict, ({k: self[k] for k in dict.keys(self)},)) + return (dict, (dict(self),)) + + # --- transparency protocol (F-TRANSPARENT): outside the observation + # machinery, the wrapper IS its plain value. Declared once here; + # boundaries consult the protocol instead of enumerating wrapper types. + __pyir_plain_type__ = dict + + def __pyir_plain_view__(self) -> dict: + # Raw snapshot (same de-watch contract as ``__copy__``): every stored + # element is choke-maintained, so raw storage IS the current value. + return dict(self) + + +class _WatchedList(list): + """Identity-adopted tracked ``list``: integer-index places with the same choke + lifecycle as :class:`_WatchedDict`; length/order changes are trace-time-only.""" + + # Slot attribute type (assigned post-construction at the adoption choke). + _pyir_label: str + + __slots__ = ("_pyir_label",) + + def _pyir_engaged(self) -> bool: + """Whether the object-side chokes fire for this access (see class doc).""" + return ( + not _WATCHED_CONTAINER_BYPASS[0] + and bool(_PYIR_SCOPE_STACK) + and _WATCHED_LIST_READ_HOOK[0] is not None + ) + + def __getitem__(self, key: Any) -> Any: + if not self._pyir_engaged() or isinstance(key, slice): + # Slice reads are trace-time snapshots (plain-list copies) -- + # declared residue. + return list.__getitem__(self, key) + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + return _WATCHED_LIST_READ_HOOK[0](self, key) + except IndexError: + # A missed lookup consumed the LENGTH (the bounds fact the trace + # acts on): record the contents snapshot, then raise (F-SPEC, + # same funnel as ``__contains__``). + _pyir_spec_record_membership(self) + raise + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + + def __contains__(self, key: Any) -> bool: + # Membership consumes the WHOLE contents through the C-level scan + # (no item choke fires): record a contents-snapshot row (F-SPEC). + if self._pyir_engaged(): + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + _pyir_spec_record_membership(self) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + return list.__contains__(self, key) + + def __setitem__(self, key: Any, value: Any) -> None: + if not self._pyir_engaged(): + list.__setitem__(self, key, value) + return + if isinstance(key, slice): + # A slice write can change the length -- structural. + self._pyir_structural("__setitem__", f"slice write [{key!r}]") + list.__setitem__(self, key, value) + return + # Diagnostics-only source attribution via the ONE sanctioned frame + # helper (never influences emission or slot identity). + filename, lineno = _first_non_dsl_caller_location() + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + value = _WATCHED_LIST_WRITE_HOOK[0]( + self, key, value, filename or "", lineno or 0 + ) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + list.__setitem__(self, key, value) + + # Structural mutators (length/order changes) are trace-time-only: refused in + # dynamic staged CF; constexpr-scope inserts freeze foreign-bound staged elements. + + def _pyir_structural(self, op_name: str, detail: str) -> None: + if self._pyir_engaged() and _WATCHED_LIST_MUTATOR_HOOK[0] is not None: + _WATCHED_CONTAINER_BYPASS[0] += 1 + try: + _WATCHED_LIST_MUTATOR_HOOK[0](self, op_name, detail) + finally: + _WATCHED_CONTAINER_BYPASS[0] -= 1 + + def append(self, value: Any) -> None: + list.append(self, value) + self._pyir_structural("append", "append()") + + def extend(self, other: Any) -> None: + list.extend(self, other) + self._pyir_structural("extend", "extend()") + + def insert(self, index: Any, value: Any) -> None: + list.insert(self, index, value) + self._pyir_structural("insert", f"insert({index!r})") + + def pop(self, *args: Any) -> Any: + # A pop reads the element's place (at the PRE-pop index) before the + # removal, so post-region consumers observe the place cell. + if self._pyir_engaged(): + idx = args[0] if args else -1 + read_val = self[idx] # __getitem__ routes the tracked read + list.pop(self, *args) + self._pyir_structural("pop", "pop()") + return read_val + val = list.pop(self, *args) + self._pyir_structural("pop", "pop()") + return val + + def remove(self, value: Any) -> None: + list.remove(self, value) + self._pyir_structural("remove", "remove()") + + def clear(self) -> None: + list.clear(self) + self._pyir_structural("clear", "clear()") + + def sort(self, *args: Any, **kwargs: Any) -> None: + list.sort(self, *args, **kwargs) + self._pyir_structural("sort", "sort()") + + def reverse(self) -> None: + list.reverse(self) + self._pyir_structural("reverse", "reverse()") + + def __delitem__(self, key: Any) -> None: + list.__delitem__(self, key) + self._pyir_structural("__delitem__", f"del [{key!r}]") + + # ``+`` keeps list.__add__ (a PLAIN list copy by design), so the augmented + # form's in-place signature intentionally diverges from it. + def __iadd__(self, other: Any) -> "_WatchedList": # type: ignore[misc] + list.extend(self, other) + self._pyir_structural("extend", "+=") + return self + + def __imul__(self, factor: Any) -> "_WatchedList": # type: ignore[misc] + list.__imul__(self, factor) + self._pyir_structural("__imul__", "*=") + return self + + # --- snapshot semantics: copies are PLAIN lists -------------------------- + + def copy(self) -> list: + return list(self) + + def __copy__(self) -> list: + return list(self) + + def __deepcopy__(self, memo: Any) -> list: + import copy as _copy + + out: list = [] + memo[id(self)] = out + for elem in list.__iter__(self): + out.append(_copy.deepcopy(elem, memo)) + return out + + def __reduce__(self) -> Any: + # Same de-watch + choke-read contract as ``_WatchedDict.__reduce__``. + if self._pyir_engaged(): + return (list, ([self[i] for i in range(list.__len__(self))],)) + return (list, (list(self),)) + + # --- transparency protocol (F-TRANSPARENT): see ``_WatchedDict``. + __pyir_plain_type__ = list + + def __pyir_plain_view__(self) -> list: + # Raw snapshot (same de-watch contract as ``__copy__``). + return list(self) + + +def _pyir_identity_opaque(value: Any) -> bool: + """Whether *value* is a wrapper the tracer mints over a tracked Python + place (watched primitive / ir-backed scalar wrapper): its OBJECT identity + does not follow the program's own object identity. Raw ``ir.Value`` + objects stay native: their stored-reference identity is a DSL-level + invariant (the no-op-cast backing share), not an A-axis place wrapper.""" + if isinstance(value, _WatchedM): + return True + from .typing import Numeric + + return isinstance(value, Numeric) and isinstance( + value.__dict__.get("value"), ir.Value + ) + + +def _pyir_judge_identity_leg(left: Any, right: Any) -> "str | None": + """Refuse-or-witness one ``is``/``is not`` leg (LangRef 3.12 section + 6.10.3): over tracker-minted wrappers the fold compares WRAPPER identity. + The fold stands where a declared fact makes it exact -- the two names + hold one object, or the tracked side's payload domain (bool/int/float) + is type-disjoint from the other operand, so ``is`` is False in both + worlds. Everything else (two distinct wrappers, or a wrapper against a + numeric primitive CPython may or may not intern) has no trace-time + answer: refuse loudly. Returns the witnessed fold kind (``None`` for an + inert leg).""" + left_opaque = _pyir_identity_opaque(left) + right_opaque = _pyir_identity_opaque(right) + if not left_opaque and not right_opaque: + return None + if left is right: + kind = "same-object" + elif left_opaque and right_opaque: + kind = None + elif type(right if left_opaque else left) in (bool, int, float): + kind = None + else: + kind = "domain-disjoint" + if kind is None: + filename, lineno = _first_non_dsl_caller_location() + opaque = left if left_opaque else right + raise DSLUserCodeError( + DiagId.PHASE_IDENTITY_ON_TRACKED, + filename=filename, + lineno=lineno, + what=type(opaque).__name__, + ) + return kind + + +def _pyir_identity_compare_choke( + left: Any, comparators: "list[Any]", ops: "list[Any]" +) -> None: + """Judge every ``is``/``is not`` leg of one comparison chain; inert + without an open PyIR trace scope.""" + if not _PYIR_SCOPE_STACK: + return + current = left + for comparator, op in zip(comparators, ops): + if op in ("is", "is not"): + _pyir_judge_identity_leg(current, comparator) + current = comparator + + +def _pyir_watched_dict_label(container: Any, key: Any) -> str: + """Diagnostic label for one adopted dict entry (``[key]``).""" + base = getattr(container, "_pyir_label", None) or "dict" + return f"{base}[{key!r}]" + + +def _pyir_watched_list_label(container: Any, key: Any) -> str: + """Diagnostic label for one adopted list element (``[i]``).""" + base = getattr(container, "_pyir_label", None) or "list" + return f"{base}[{key!r}]" + + +def _pyir_adopt_container_value( + owner: Any, + slot_name: Any, + obj: Any, + label: "str | None", + raw_type: type, + watched_cls: type, + fill: Any, + entries: Any, +) -> Any: + """Shared skeleton of the dict/list adoption twins: watched passthrough, + exact-type gate, unrepointable-holder guard, adopt-once memo (aliases + converge) with recursive child adoption, holder re-point. Returns the + watched twin, or *obj* untouched on a pass-through path (the callers + skip their container-specific tails then).""" + if isinstance(obj, watched_cls): + w = obj + elif type(obj) is not raw_type: + return obj + elif ( + owner is not None + and slot_name is not None + and not _pyir_holder_slot_refers(owner, slot_name, obj) + ): + # The holder in hand cannot be re-pointed (e.g. a class attribute): + # adopting would mint a stale watched twin -- keep the raw object. + return obj + else: + w = _WATCHED_CONTAINER_ADOPTIONS.get(id(obj)) + if w is None: + w = watched_cls() + w._pyir_label = str(label) if label is not None else raw_type.__name__ + fill(w, obj) + _WATCHED_CONTAINER_ADOPTIONS[id(obj)] = w + _WATCHED_CONTAINER_KEEPALIVE.append(obj) + for k, v in entries(w): + if type(v) is dict or isinstance(v, _WatchedDict): + _pyir_adopt_dict_value(w, k, v, label=f"{w._pyir_label}[{k!r}]") + elif type(v) is list or isinstance(v, _WatchedList): + _pyir_adopt_list_value(w, k, v, label=f"{w._pyir_label}[{k!r}]") + if w is not obj and owner is not None and slot_name is not None: + _pyir_replace_holder_slot(owner, slot_name, obj, w) + return w + + +def _pyir_adopt_dict_value( + owner: Any, slot_name: Any, d: Any, label: "str | None" = None +) -> Any: + """Adopt plain dict *d* as its place-owning :class:`_WatchedDict`: adopt-once + (aliases converge), known holder slots re-pointed; non-exact dicts pass through.""" + w = _pyir_adopt_container_value( + owner, + slot_name, + d, + label, + dict, + _WatchedDict, + dict.update, + lambda w: [(k, dict.__getitem__(w, k)) for k in list(dict.keys(w))], + ) + if not isinstance(w, _WatchedDict): + return w # pass-through: nothing was adopted + # Instance-storage fact: a mapping adopted through the ``__dict__`` + # pseudo-slot IS the owner's attribute storage (LangRef 3.3.2), so its + # str-keyed item places unify with the owner's attr places. + if ( + slot_name == "__dict__" + and owner is not None + and not isinstance(owner, type) + and getattr(owner, "__dict__", None) is w + ): + try: + w._pyir_instance_of = _weakref.ref(owner) + except TypeError: + w._pyir_instance_of = owner # non-weakrefable: trace-scoped strong ref + _pyir_spec_chain_container(w, owner, slot_name) + return w + + +def _pyir_instance_dict_owner(container: Any, key: Any) -> Any: + """The instance whose attribute storage *container* IS (recorded at the + ``vars(obj)``/``obj.__dict__`` adoption): a str-keyed item access names the + SAME place as the attr spelling, so the item chokes route it to the + instance's attr row -- one place row per storage. ``None`` when the fact + is absent, expired (storage re-pointed), the key is not a str, or a data + descriptor on the class intercepts the attr spelling (different storage).""" + ref = getattr(container, "_pyir_instance_of", None) + if ref is None or not isinstance(key, str): + return None + owner = ref() if isinstance(ref, _weakref.ref) else ref + if owner is None: + return None + try: + if getattr(owner, "__dict__", None) is not container: + return None + for klass in type(owner).__mro__: + desc = klass.__dict__.get(key) + if desc is None: + continue + if hasattr(type(desc), "__set__") or hasattr(type(desc), "__delete__"): + return None + break + except Exception: + return None + return owner + + +def _pyir_spec_chain_container(twin: Any, owner: Any, slot_name: Any) -> None: + """F-SPEC root chaining: an adopted container IS the persistent value of its + holder slot, so its token roots at the holder's root plus that hop + (first-wins; owner-less adoptions stay trace-internal).""" + try: + if not isinstance(twin, (_WatchedDict, _WatchedList)): + return + _pyir_spec_chain_value(twin, owner, slot_name) + except Exception: + pass # chaining must never break the adoption it observes + + +def _pyir_spec_value_is_ir_wrapper(value: Any) -> bool: + """True for a value of the MLIR value-tree wrapper family (staged, or its + class implements the ``__extract_mlir_values__`` protocol): such an object + observed at a persistent place during the trace is TRACE-EPOCH -- its + interior is derived from this trace's IR, so a spec-row chain through it + is trace-internal, never launch-recheckable.""" + if _is_staged_value(value): + return True + return getattr(type(value), "__extract_mlir_values__", None) is not None + + +def _pyir_spec_chain_value( + value: Any, owner: Any, slot_name: Any = None, steps: "tuple | None" = None +) -> None: + """F-SPEC root chaining for ANY owner-slot read hop: the value observed at a + ``(owner, slot)`` place roots at the owner's root plus that hop, so scalar + legs one composition deeper resolve a re-derivable root path (first-wins; + hops off unrooted owners stay trace-internal until the owner roots). + A hop through a TRACE-EPOCH object never chains: re-resolving it at launch + would walk this trace's consumed IR (the value stamps trace-born instead, + so deeper hops stay trace-internal too).""" + try: + if value is None or owner is None: + return + if _is_staged_value(value) or isinstance(value, _WatchedM): + return + if type(value) in (bool, int, float, str, bytes) or isinstance( + value, _enum.Enum + ): + return + if isinstance(value, (types.ModuleType, type)): + return # symbol bases carry canonical roots, never holder chains + if _pyir_spec_value_is_ir_wrapper(value): + _pyir_spec_stamp_trace_born(value) + return + o_tok_early = _owner_token(owner) + if o_tok_early is not None and o_tok_early in _PYIR_SPEC_TRACE_BORN_TOKENS: + _pyir_spec_stamp_trace_born(value) + return + if steps is None: + if slot_name is None: + return + if isinstance(slot_name, _PlaceSeg): + steps = slot_name.steps() + elif isinstance(owner, (dict, list, tuple)) or not isinstance( + slot_name, str + ): + steps = (("item", slot_name),) + else: + steps = (("attr", slot_name),) + v_tok = _owner_token(value) + o_tok = _owner_token(owner) + if v_tok is None or o_tok is None or v_tok == o_tok: + return + _PYIR_SPEC_CONTAINER_CHAIN.setdefault(v_tok, (o_tok, tuple(steps))) + except Exception: + pass # chaining must never break the read it observes + + +def _pyir_spec_stamp_trace_born(value: Any) -> None: + """Stamp *value* TRACE-EPOCH (DF-3): its rows are trace-internal and never + record. The strong pin makes the id-token provably name this object for + the whole trace (a dead stamped wrapper can never alias a fresh holder).""" + tok = _owner_token(value) + if tok is not None: + _PYIR_SPEC_TRACE_BORN_TOKENS.add(tok) + _PYIR_SPEC_TRACE_BORN_PINS.append(value) + + +def _pyir_holder_slot_refers(owner: Any, slot_name: Any, obj: Any) -> bool: + """Whether the ``(owner, slot_name)`` slot references *obj* through re-pointable + storage; adopting through anything else would split a stale watched twin. + ``__dict__`` is the instance-storage pseudo-slot (the mapping IS the + owner's assignable storage, so it is holder storage by definition).""" + try: + if isinstance(owner, dict): + return dict.__getitem__(owner, slot_name) is obj + if isinstance(owner, list): + return list.__getitem__(owner, slot_name) is obj + if slot_name == "__dict__" and not isinstance(owner, type): + return getattr(owner, "__dict__", None) is obj + storage = _instance_storage_items(owner) + return storage is not None and storage.get(slot_name) is obj + except Exception: + return False + + +def _pyir_replace_holder_slot(owner: Any, slot_name: Any, old: Any, new: Any) -> None: + """Re-point the ``(owner, slot_name)`` slot from *old* to *new*, bypassing user + hooks and frozen guards; best-effort no-op when the slot moved on.""" + try: + if isinstance(owner, dict): + if dict.__getitem__(owner, slot_name) is old: + dict.__setitem__(owner, slot_name, new) + elif isinstance(owner, list): + if list.__getitem__(owner, slot_name) is old: + list.__setitem__(owner, slot_name, new) + elif slot_name == "__dict__" and not isinstance(owner, type): + if getattr(owner, "__dict__", None) is old: + _pyir_setattr_raw(owner, "__dict__", new) + else: + storage = _instance_storage_items(owner) + if storage is not None and storage.get(slot_name) is old: + _pyir_setattr_raw(owner, slot_name, new) + except Exception: + pass + + +def _pyir_adopt_list_value( + owner: Any, slot_name: Any, l: Any, label: "str | None" = None +) -> Any: + """Adopt plain list *l* as its place-owning :class:`_WatchedList` (exact list + sibling of :func:`_pyir_adopt_dict_value`); non-exact lists pass through.""" + w = _pyir_adopt_container_value( + owner, + slot_name, + l, + label, + list, + _WatchedList, + list.extend, + lambda w: [(i, list.__getitem__(w, i)) for i in range(list.__len__(w))], + ) + if not isinstance(w, _WatchedList): + return w # pass-through: nothing was adopted + _pyir_spec_chain_container(w, owner, slot_name) + return w + + +def _pyir_register_container_held_object(obj: Any) -> None: + """Record *obj* as a tracked-container leg value (the root -> container-leg -> + attr-leg composition fact); value-tree-protocol objects are skipped.""" + if obj is None or isinstance( + obj, + ( + int, + float, + bool, + str, + bytes, + type, + types.ModuleType, + ir.Value, + dict, + list, + tuple, + set, + frozenset, + ), + ): + return + oid = id(obj) + if oid in _WATCHED_CONTAINER_HELD_OBJECTS: + return + if _is_staged_value(obj) or _implements_dynamic_expression(obj): + return + if not _has_instance_storage(obj): + return + _WATCHED_CONTAINER_HELD_OBJECTS.add(oid) + _WATCHED_CONTAINER_HELD_KEEPALIVE.append(obj) + + +def _pyir_owner_is_container_held(owner: Any) -> bool: + """Whether *owner* was discovered as a tracked-container leg value this trace + (qualifies its attr legs for the adopted meta lifecycle).""" + return owner is not None and id(owner) in _WATCHED_CONTAINER_HELD_OBJECTS + + +def _pyir_adopt_sequence_elems(seq: Any, prefix: str) -> None: + """Element walk shared by the two tuple/list arms of + :func:`_pyir_adopt_containers_under`: dict/list elements adopt in place + (a tuple slot cannot re-point, so its container legs keep the raw value), + the item hop declares each element's root either way, and opaque + elements recurse.""" + for i, elem in enumerate(seq): + if type(elem) is dict and isinstance(seq, list): + seq[i] = _pyir_adopt_dict_value(None, None, elem, label=f"{prefix}[{i}]") + _pyir_spec_chain_value(seq[i], seq, steps=(("item", i),)) + elif type(elem) is list and isinstance(seq, list): + seq[i] = _pyir_adopt_list_value(None, None, elem, label=f"{prefix}[{i}]") + _pyir_spec_chain_value(seq[i], seq, steps=(("item", i),)) + else: + _pyir_spec_chain_value(elem, seq, steps=(("item", i),)) + _pyir_adopt_containers_under(elem) + + +# Element types the sequence-adoption walk provably does nothing with; exact +# types only, mirroring the two callees' own early-return checks. +_SEQ_WALK_INERT_PRIMS = frozenset({int, float, bool, str, bytes, type(None)}) + + +def _mark_container_walk_visited(obj: Any, oid: int) -> None: + """Memoize *obj* as adoption-walked. The visited object is pinned for the + trace: a GC'd-and-recycled id must never make a fresh object look + already-walked (its containers would silently escape adoption). + + The pin is why only objects whose re-walk could change something are + memoized: it also keeps the object's candidate-holder registry row alive, + and every staged region sweeps that registry, so pinning each staged scalar + sighted at a ``@dsl_user_op`` boundary would make the per-region sweep grow + with trace history instead of with live captured state.""" + _WATCHED_CONTAINER_WALKED.add(oid) + _WATCHED_CONTAINER_WALKED_KEEPALIVE.append(obj) + + +def _sequence_walk_is_inert(seq: Any) -> bool: + """True when re-walking *seq* provably changes nothing, so the memo (and the + trace-long pin it needs) buys nothing. + + Inert = every element is a shape BOTH halves of + :func:`_pyir_adopt_sequence_elems` return on before any recorded effect: a + bare Python scalar (``_pyir_spec_chain_value`` line 1 / the adoption walk's + head filter) or a staged value-tree object (chaining returns on staged, + adoption returns on the value-tree protocol). Anything else -- a raw + ``dict``/``list`` element, an opaque object, a bare ``ir.Value`` (which + stamps trace-born) -- keeps the memo, since re-walking it would repeat a + recorded effect.""" + for elem in seq: + if type(elem) in _SEQ_WALK_INERT_PRIMS: + continue + if _is_staged_value(elem) and _implements_dynamic_expression(elem): + continue + return False + return True + + +def _pyir_adopt_containers_under(obj: Any) -> None: + """Adoption walk: replace plain dict/list fields reachable from *obj* with + watched instances; memoized per trace, value-tree objects and bare roots skipped.""" + if obj is None or isinstance( + obj, + (int, float, bool, str, bytes, type, types.ModuleType, ir.Value, dict), + ): + return + oid = id(obj) + if oid in _WATCHED_CONTAINER_WALKED: + return + if isinstance(obj, (tuple, list)): + if not _sequence_walk_is_inert(obj): + _mark_container_walk_visited(obj, oid) + _pyir_adopt_sequence_elems(obj, prefix="") + return + if _implements_dynamic_expression(obj): + return # nothing to walk, so nothing to memoize + storage = _instance_storage_items(obj) + if storage is None: + return + _mark_container_walk_visited(obj, oid) + for attr, val in list(storage.items()): + if isinstance(attr, str) and attr.startswith("__"): + continue + if type(val) is dict: + _pyir_adopt_dict_value(obj, attr, val, label=attr) + elif type(val) is list: + _pyir_adopt_list_value(obj, attr, val, label=attr) + elif isinstance(val, (_WatchedDict, _WatchedList)): + # A twin persisting from an earlier trace: re-declare its root + # chain (the chain registry is trace-scoped, the twin is not). + _pyir_spec_chain_container(val, obj, attr) + continue + elif isinstance(val, (tuple, list)): + _pyir_spec_chain_value(val, obj, attr) + _pyir_adopt_sequence_elems(val, prefix=str(attr)) + elif not isinstance(val, (int, float, bool, str, bytes)): + _pyir_adopt_containers_under(val) + + +def _pyir_holder_read(holder: Any, attr: Any) -> Any: + """Raw read of a ``(holder, attr)`` slot -- the mirror of + :func:`_pyir_holder_store` (bypasses watched chokes, user hooks, and + descriptors: instance/class STORAGE only).""" + if isinstance(holder, dict): + return dict.__getitem__(holder, attr) + if isinstance(holder, list): + return list.__getitem__(holder, attr) + if isinstance(holder, types.CellType): + return holder.cell_contents + if isinstance(holder, type): + return holder.__dict__[attr] + storage = _instance_storage_items(holder) + if storage is None or attr not in storage: + raise AttributeError(attr) + return storage[attr] + + +def _pyir_record_host_restore(owner: Any, slot_name: Any, pre_value: Any) -> None: + """Witness a host place's last scalar meta binding before a write; trace + close restores it if the place is left holding this trace's wrapper.""" + if owner is None or slot_name is None: + return + if isinstance(pre_value, _WatchedM): + pre_value = pre_value._pyir_raw_payload + if type(pre_value) not in (bool, int, float): + return # non-scalar pre-binding: a leftover keeps the stale-epoch refusal + try: + ctx_id = id(ir.Context.current) + except Exception: + return + try: + _PYIR_HOST_RESTORE[(id(owner), slot_name)] = ( + owner, + slot_name, + pre_value, + ctx_id, + ) + except TypeError: + pass # unhashable slot key: a leftover keeps the stale-epoch refusal + + +def _pyir_restore_host_places() -> None: + """Trace close: re-point every recorded host place still bound to a wrapper + THIS compilation minted back to its pre-staging meta binding (the documented + post-trace state of a promoted place), and re-bake its F-SPEC rows so a + no-retrace reuse verifies the restored state.""" + for owner, slot_name, pre_value, ctx_id in list(_PYIR_HOST_RESTORE.values()): + try: + cur = _pyir_holder_read(owner, slot_name) + except Exception: + continue + if getattr(cur, "_pyir_birth_ctx", None) != ctx_id: + continue # meta rebinds and adopted twins stay: last host binding wins + try: + _pyir_holder_store(owner, slot_name, pre_value) + except Exception: + continue + _pyir_spec_record_write(_make_slot_key(None, owner, slot_name), pre_value) + _PYIR_HOST_RESTORE.clear() + + +def _pyir_holder_store(holder: Any, attr: Any, value: Any) -> None: + """Raw write of a ``(holder, attr)`` slot (dict entry, list element, or + attribute), bypassing watched chokes, user hooks, and frozen guards.""" + try: + _pyir_record_host_restore(holder, attr, _pyir_holder_read(holder, attr)) + except Exception: + pass # no pre-binding (first def): nothing to restore + if isinstance(holder, dict): + dict.__setitem__(holder, attr, value) + elif isinstance(holder, list): + list.__setitem__(holder, attr, value) + elif isinstance(holder, types.CellType): + # Closure cell holder: the single ``cell_contents`` binding (the + # *attr* carries the closure variable name for diagnostics only). + holder.cell_contents = value + elif isinstance(holder, type): + # Class-attribute holder (a context-manager counter): + # ``object.__setattr__`` rejects classes, plain ``setattr`` works. + setattr(holder, attr, value) + else: + _pyir_setattr_raw(holder, attr, value) + + +def _replace_value_uses(old_val: "ir.Value", new_val: "ir.Value") -> bool: + """Best-effort ``replaceAllUsesWith`` across MLIR Python binding versions.""" + for method_name in ("replace_all_uses_with", "replaceAllUsesWith"): + fn = getattr(old_val, method_name, None) + if fn is not None: + try: + fn(new_val) + return True + except Exception: + continue + return False + + +def _forget_region_local_slot_state(owner: Any, slot_name: Any) -> None: + """Drop a region-local meta slot's baked-constant bookkeeping so a sibling + region starts from op-entry state; slots already in ``_slot_refs`` untouched.""" + slot_key = _make_slot_key(None, owner, slot_name) + if slot_key is None or slot_key in _slot_refs: + return + _meta_uses.pop(slot_key, None) + _slot_first_def_inside_cf.pop(slot_key, None) + _slot_first_def_block.pop(slot_key, None) + _slot_binding_depth.pop(slot_key, None) + + +def _exit_function_trace() -> None: + """Clear D1 per-trace state. Called from ``_jit_scope.finally``.""" + # Host restore first: user-visible bindings must not keep this trace's + # staged wrappers past compilation (documented promoted-place semantics). + _pyir_restore_host_places() + # Scratch-admission enforcement runs at trace close, while the emitted + # ops are live: every consuming read of an admitted cell needs an + # in-window reaching store. A trace already failing keeps its own error; + # a sweep refusal is held so the state clearing below still runs. + _scratch_refusal: "BaseException | None" = None + try: + if _sys.exc_info()[0] is None: + _pyir_verify_scratch_admissions() + except DSLUserCodeError as _e: + _scratch_refusal = _e + finally: + _PYIR_SCRATCH_ADMISSIONS.clear() + _PYIR_MEMREF_LAST_STORE.clear() + # F-SPEC seal: snapshot the specialization ledger of the trace that just + # closed; the compile flow moves it onto the JitCompiledFunction. An + # empty-and-complete ledger seals as None (nothing to validate). + _PYIR_SPEC_SEALED[0] = ( + ( + dict(_PYIR_SPEC_RECORD), + _PYIR_SPEC_COMPLETE[0], + _PYIR_SPEC_ENTRY_FUNC[0], + _pyir_spec_sealed_receiver(), + _PYIR_SPEC_ENTRY_SIG[0], + ) + if (_PYIR_SPEC_RECORD or not _PYIR_SPEC_COMPLETE[0]) + else None + ) + _reset_spec_state() + _meta_uses.clear() + _meta_idempotent_write_anchors.clear() + _slot_refs.clear() + _slot_templates.clear() + # Ref write-epochs / promotion rewrite map key MLIR-context-bound + # ir.Values; never cross traces. + _REF_WRITE_EPOCH.clear() + _META_CONST_REPLACEMENTS.clear() + _PYIR_BOUNDARY_META_CELL_READS.clear() + _PYIR_BOUNDARY_FLIP_GUARD_READS.clear() + _PYIR_HOST_RESTORE.clear() + _PYIR_STRUCTURAL_META_CONSUMPTIONS.clear() + _PYIR_STRUCTURAL_META_CONSUMPTION_SITES.clear() + _PYIR_ARM_LOCAL_META_WRITES.clear() + _PYIR_STAGED_LITERAL_FOLD_WITNESSES.clear() + _PYIR_STAGED_LITERAL_FOLD_KEEPALIVE.clear() + _PYIR_STAGED_HASH_WITNESS[0] = 0 + _slot_first_def_inside_cf.clear() + _slot_first_def_block.clear() + _slot_first_def_block_any.clear() + _slot_first_def_depth.clear() + _slot_first_def_depth_any.clear() + _slot_binding_depth.clear() + _pyir_open_loop_body_blocks.clear() + # Clear the slot registry; it holds MLIR-context-bound MutableValue + # instances which must not survive across traces. + _SLOT_REGISTRY.clear() + _OWNER_KEEPALIVE.clear() + _PYIR_SLOT_HOLDERS.clear() + _PYIR_PROMOTED_PLACE_LEAVES.clear() + # Candidate-holder registry is trace-scoped (its strong non-weakrefable + # entries must not outlive the trace). Its id-keyed stamp/segment side + # rows go with it: clearing the registry discards the weakrefs whose + # callbacks would otherwise pop them at object death. + _PYIR_CANDIDATE_HOLDERS.clear() + _PYIR_HOLDER_WRITE_STAMPS.clear() + _PYIR_GATHER_SEGMENTS.clear() + _PYIR_SIGHTING_SETTLED.clear() + _PYIR_WRITE_CLOCK[0] = 0 + _PYIR_CANDIDATE_REGISTRY_ACTIVE[0] = False + # Watched-container adoption state is trace-scoped; adopted instances stay + # embedded in user objects, inert without an open PyIR scope. + _WATCHED_CONTAINER_ADOPTIONS.clear() + # Armed opaque-owner audits hold strong owner references; trace-scoped. + _PYIR_OPAQUE_OWNER_AUDITS.clear() + # In-CF key-creation records hold strong container references; trace-scoped. + _PYIR_DICT_CF_CREATED_KEYS.clear() + _WATCHED_CONTAINER_KEEPALIVE.clear() + _PYIR_SPEC_CONTAINER_CHAIN.clear() + _WATCHED_CONTAINER_WALKED.clear() + _WATCHED_CONTAINER_WALKED_KEEPALIVE.clear() + _WATCHED_CONTAINER_BYPASS[0] = 0 + _WATCHED_CONTAINER_HELD_OBJECTS.clear() + _WATCHED_CONTAINER_HELD_KEEPALIVE.clear() + # Superseded-generation detector state: stamp records hold MutableValues + # (MLIR-context-bound) and id-keyed entries; none may survive the trace. + _PYIR_CF_ATTR_FIRST_DEFS.clear() + _SUPERSEDED_GENERATIONS.clear() + _SUPERSEDED_LEAF_WRAPPERS.clear() + _GEN_OBJ_KEEPALIVE.clear() + _BINDING_GEN.clear() + _BINDING_REBIND_CELLS.clear() + _ALIAS_CAPTURES.clear() + _ALIAS_CAPTURE_ROOT_NAMES.clear() + _COMPOUND_BINDING_SLOTS.clear() + _PYIR_GEN_EVENT_SUPPRESS[0] = 0 + # Place layer: clear the per-trace scope stack, owner tokens, and place + # registry (holds MLIR-context-bound MutableValues). + _reset_scope_state() + if _scratch_refusal is not None: + raise _scratch_refusal + + +def _pyir_take_sealed_spec() -> "tuple | None": + """Move the last sealed F-SPEC snapshot to the caller (one collection per + trace close); ``None`` when no PyIR trace sealed since the last take.""" + sealed = _PYIR_SPEC_SEALED[0] + _PYIR_SPEC_SEALED[0] = None + return sealed + + +def _pyir_assert_entry_attested(func_body: Any, dsl_obj: Any) -> None: + """V-9 ATTESTED: a trace entry must carry the current rewrite stamp or a + declared-native mark; anything else traces with invisible reads/writes.""" + if not is_pyir_enabled(): + return + ver = getattr(getattr(dsl_obj, "preprocessor", None), "choke_set_version", None) + if ver is None or not getattr(dsl_obj, "enable_preprocessor", True): + return # the DSL declared a whole-hog native (non-PyIR) trace mode + fn = getattr(func_body, "__func__", func_body) + if getattr(fn, "__pyir_rewritten__", None) == ver: + return + if ( + getattr(fn, "__pyir_native__", False) + or getattr(fn, "_preprocess_enabled", True) is False + ): + return + raise DSLUserCodeError( + DiagId.INSTRUMENTATION_GAP, + name=getattr(fn, "__qualname__", getattr(fn, "__name__", repr(fn))), + ) + + +def _slot_store_for_tier1( + owner: object, *, create: bool = False +) -> "dict[Any, MutableValue] | None": + """Return the tier-1 ``__pyir_slots__`` dict on *owner* (``None`` without a + writable ``__dict__``); *create* installs one bypassing frozen setattr.""" + owner_dict = getattr(owner, "__dict__", None) + if owner_dict is None: + return None + store = owner_dict.get(_SLOT_STORE_ATTR) + if store is None: + if not create: + return None + store = {} + try: + _pyir_setattr_raw(owner, _SLOT_STORE_ATTR, store) + except (AttributeError, TypeError): + return None + return store + + +def _slot_storage_available(owner: Any) -> bool: + """True if *owner* can host slot storage; every non-``None`` owner has some + storage path (tier-1 ``__dict__`` or ``_SLOT_REGISTRY``).""" + return owner is not None + + +def _registry_owner(owner: Any) -> bool: + """True if *owner*'s slots route through ``_SLOT_REGISTRY`` (any owner + without a WRITABLE instance ``__dict__``: none at all, or a read-only + mapping like a class's mappingproxy); writable-dict owners use tier-1.""" + if owner is None: + return False + d = getattr(owner, "__dict__", None) + return d is None or not isinstance(d, dict) + + +def _get_slot_mv(owner: Any, slot_name: Any) -> "MutableValue | None": + """Look up the ``MutableValue`` bound to ``(owner, slot_name)`` (registry or + tier-1); ``owner=None`` resolves the ownerless local place directly.""" + if owner is None: + # Ownerless LOCAL place: place registry only, on the F-CEPLACE + # instance-qualified key (constexpr scopes never gate production). + if slot_name is None: + return None + try: + place = _make_slot_key(slot_name, None, None) + except Exception: + return None + return _PLACE_REGISTRY.get(place) if place is not None else None + # Id-keyed resolution is primary: registry slots for ownerless-storage + # owners, the tier-1 ``__pyir_slots__`` dict otherwise. + if _registry_owner(owner): + slot_id = _make_slot_id(owner, slot_name) + live_mv = _SLOT_REGISTRY.get(slot_id) + else: + live_mv = None + store = _slot_store_for_tier1(owner) + if store is not None: + live_mv = store.get(slot_name) + # On an id-key MISS resolve by PLACE (a reconstructed holder keeps its + # token-rooted place); an attr place is unroll-invariant, so the place + # layer runs inside constexpr scopes too. + place_fallback = None + _sc_place = None + _sc_place_mv = None + try: + _place = _corrected_place_for(owner, slot_name) + _place_mv = _PLACE_REGISTRY.get(_place) if _place is not None else None + _sc_place, _sc_place_mv = _place, _place_mv + if live_mv is None: + place_fallback = _place_mv + except Exception: + pass + # Emission self-check (one-cell-per-place): the diagnostic raise must + # propagate; pre-filtered on cell identity. + if _sc_place_mv is not None and live_mv is not None and _sc_place_mv is not live_mv: + _pyir_emission_self_check( + "read", owner, slot_name, live_mv, _sc_place, _sc_place_mv + ) + if live_mv is None: + return place_fallback + return live_mv + + +def _pyir_track_slot_holder(owner: Any) -> None: + """Record *owner* in the slot-holder registry for the region-close carry sweep; + id-keyed, never invokes user ``__hash__``/``__eq__``.""" + oid = id(owner) + if oid not in _PYIR_SLOT_HOLDERS: + try: + _PYIR_SLOT_HOLDERS[oid] = _weakref.ref( + owner, lambda _r: _PYIR_SLOT_HOLDERS.pop(oid, None) + ) + except TypeError: + _PYIR_SLOT_HOLDERS[oid] = owner + # A slot owner is also a candidate root for the captured-holder gathers. + _pyir_register_candidate_holder(owner) + + +def _pyir_refuse_del_finalizer_owner(obj: Any) -> None: + """S4 floor (LangRef 3.12 section 3.3.1): refuse admitting an instance of + a user class defining ``__del__`` into the ledger's identity domain -- + its finalizer runs at a GC-determined instant, which has no binding + position in the traced program.""" + from . import pyir_class_facts as _pcf + + definer = _pcf.del_finalizer_definer(type(obj)) + if definer is None: + return + from .pyir_call_boundary import _pyir_boundary_module_is_user + + if not _pyir_boundary_module_is_user(getattr(definer, "__module__", None)): + return # DSL/stdlib finalizers stay under the wrapper-consumer contract + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.OWNER_DEL_FINALIZER_AT_BIRTH, + filename=filename, + lineno=lineno, + cls=type(obj).__name__, + definer=definer.__qualname__, + ) + + +# Per-class verdict for the sighting front gate: True when instances of the +# class provably no-op through BOTH the adoption walk (head filter) and the +# registry filters, so a sighting returns before either. Built from class-shape +# facts only (layout filters, MRO-declared protocol, dynamic attribute hooks), +# so the verdict is stable; same doctrine as ``_VT_PROTOCOL_CLASS_CACHE``. +_PYIR_SIGHTING_INERT_CLASSES: "dict[type, bool]" = {} + + +def _pyir_sighting_inert_class(cls: type) -> bool: + """Whether instances of *cls* provably no-op through a candidate sighting.""" + if cls is type(None): + return True + if issubclass(cls, (tuple, list)): + return False # the adoption walk recurses into sequence elements + if issubclass(cls, (int, float, bool, str, bytes, type, types.ModuleType, dict)): + return True # head-filtered by both the adoption walk and the registry + if issubclass(cls, ir.Value): + # A bare Value registers only when it implements the value-tree + # protocol. The protocol is class-declared (grep-gated, see + # ``_vt_protocol_class_bucket``), so with default attribute access and + # no declaration in the MRO the instance check can never turn it on. + declared = any( + "__extract_mlir_values__" in k.__dict__ for k in cls.__mro__ + ) and any("__new_from_mlir_values__" in k.__dict__ for k in cls.__mro__) + if declared: + return False + return ( + cls.__getattribute__ is object.__getattribute__ + and getattr(cls, "__getattr__", None) is None + ) + return False + + +def _pyir_register_candidate_holder(obj: Any) -> None: + """Register *obj* as a candidate captured-holder walk root; admits exactly the + shapes the walkers can yield holders from, id-deduped in sighting order. + + Front gates: a sighting of an inert-class instance (provably filtered by + type alone) or of a settled id (registered, adoption settled) returns + before the adoption walk; both verdicts assert the whole body is a no-op.""" + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + cls = type(obj) + inert = _PYIR_SIGHTING_INERT_CLASSES.get(cls) + if inert is None: + inert = _PYIR_SIGHTING_INERT_CLASSES[cls] = _pyir_sighting_inert_class(cls) + if inert: + return + oid = id(obj) + if oid in _PYIR_SIGHTING_SETTLED: + return + # Adoption runs at every candidate sighting, memoized per object; placed + # before the container filter so tuple/list ELEMENTS still adopt. + adopted = _WATCHED_DICT_READ_HOOK[0] is not None + if adopted: + try: + _pyir_adopt_containers_under(obj) + except Exception: + pass + if obj is None or isinstance( + obj, (int, float, bool, str, bytes, type, types.ModuleType, tuple, list, dict) + ): + return + if oid in _PYIR_CANDIDATE_HOLDERS: + _pyir_settle_sighting(obj, oid, adopted) + return + if isinstance(obj, ir.Value) and not _implements_dynamic_expression(obj): + return + if not _has_instance_storage(obj): + return + # S4 floor at ledger admission: a user-defined `__del__` fires at a + # GC-determined instant with no binding position; judged BEFORE the memo + # insert so an intake-swallowed refusal re-fires at the next sighting. + _pyir_refuse_del_finalizer_owner(obj) + try: + _PYIR_CANDIDATE_HOLDERS[oid] = _weakref.ref(obj, _pyir_drop_candidate_rows(oid)) + # Write-stamp row minted with the registry entry (and dying with it), + # so a stamp/segment row can never describe a reused id. + _PYIR_HOLDER_WRITE_STAMPS[oid] = 0 + _PYIR_GATHER_SEGMENTS.pop(oid, None) + except TypeError: + # Non-weakrefable but ``__dict__``-backed: keep a strong, trace-scoped + # reference so the sweep can still reach it (cleared at trace exit). + _PYIR_CANDIDATE_HOLDERS[oid] = obj + _pyir_settle_sighting(obj, oid, adopted) + + +def _pyir_settle_sighting(obj: Any, oid: int, adopted: bool) -> None: + """Mark a REGISTERED *obj* settled when its re-sighting is provably a + no-op: the registry insert dedups on the id, and the adoption walk either + memoized it or head-filters it by the value-tree protocol. Skipped while + the adoption hook is down, since a later arming would still need the walk.""" + if adopted and ( + oid in _WATCHED_CONTAINER_WALKED or _implements_dynamic_expression(obj) + ): + _PYIR_SIGHTING_SETTLED.add(oid) + + +def _pyir_drop_candidate_rows(oid: int) -> "Callable[[Any], None]": + """Registry weakref callback dropping every id-keyed candidate row.""" + + def _drop(_r: Any) -> None: + _PYIR_CANDIDATE_HOLDERS.pop(oid, None) + _PYIR_HOLDER_WRITE_STAMPS.pop(oid, None) + _PYIR_GATHER_SEGMENTS.pop(oid, None) + _PYIR_SIGHTING_SETTLED.discard(oid) + + return _drop + + +def _pyir_register_trace_args( + args: Any, kwargs: Any = None, sig: Any = None, entry_func: Any = None +) -> None: + """Trace-arg intake: register each argument (tuple/list unpacked one level) + as a candidate holder root. Never raises. + + An OUTER intake (``entry_func`` given) also opens the F-SPEC ledger for + this trace: argument objects become token roots, entry parameter names + become pending local roots, and the entry function's own closure cells + become closure roots.""" + try: + if not is_pyir_enabled(): + return + # Arm the fast registration gate for this (and any nested) trace; the + # per-sighting sites then pay one flag test instead of an env consult. + _PYIR_CANDIDATE_REGISTRY_ACTIVE[0] = True + for a in args or (): + _pyir_register_candidate_holder(a) + if isinstance(a, (tuple, list)): + for e in a: + _pyir_register_candidate_holder(e) + if kwargs: + for v in kwargs.values(): + _pyir_register_candidate_holder(v) + if isinstance(v, (tuple, list)): + for e in v: + _pyir_register_candidate_holder(e) + if entry_func is not None: + _pyir_spec_open_trace(args, kwargs, sig, entry_func) + except Exception: + pass + + +# --- F-SPEC producers: root registration + the one recording funnel -------- + + +def _pyir_spec_sealed_receiver() -> Any: + """The receiver to seal with the record: the live first-binding object, + kept ONLY when the record actually roots rows at the receiver parameter + (nothing else needs it, and holding an unrelated first argument would pin + its lifetime). A receiver that died before the seal leaves rooted rows + with no live verification root: seal the DEAD sentinel so re-entry + refuses instead of skipping.""" + cand = _PYIR_SPEC_ENTRY_RECEIVER[0] + if cand is None: + return None + ref, first_name = cand + if not any(rp[0] == "arg" and rp[1] == first_name for rp in _PYIR_SPEC_RECORD): + return None + live = ref() + return live if live is not None else _SPEC_RECEIVER_DEAD + + +def _pyir_spec_signature_facts(fn: Any, sig: Any) -> "tuple | None": + """Entry-signature facts captured once at trace open and sealed with the + record (DF-1): ``(first_param_name, {name: ("pos", i) | ("kwonly",)})``. + Launch resolution then reads ``__defaults__[i]`` / ``__kwdefaults__[name]`` + -- pure data reads that still see a REASSIGNED default (drift preserved) + -- and never re-derives call structure by executing signature code. The + index map is validated against the live ``__defaults__``/``__kwdefaults__`` + shape here; on disagreement (a decorated entry whose signature lies) the + map seals ``None`` and default-rooted rows become unverifiable.""" + if sig is None: + return None + try: + params = list(sig.parameters.values()) + first_param = params[0].name if params else None + pos_defaulted = [ + p + for p in params + if p.kind + in ( + inspect.Parameter.POSITIONAL_ONLY, + inspect.Parameter.POSITIONAL_OR_KEYWORD, + ) + and p.default is not inspect.Parameter.empty + ] + kwonly_defaulted = [ + p + for p in params + if p.kind is inspect.Parameter.KEYWORD_ONLY + and p.default is not inspect.Parameter.empty + ] + live_defaults = getattr(fn, "__defaults__", None) or () + live_kwdefaults = getattr(fn, "__kwdefaults__", None) or {} + consistent = len(live_defaults) == len(pos_defaulted) and all( + live_defaults[i] is p.default for i, p in enumerate(pos_defaulted) + ) + consistent = consistent and all( + p.name in live_kwdefaults and live_kwdefaults[p.name] is p.default + for p in kwonly_defaulted + ) + if not consistent: + return (first_param, None) + default_index: "dict[str, tuple]" = { + p.name: ("pos", i) for i, p in enumerate(pos_defaulted) + } + for p in kwonly_defaulted: + default_index[p.name] = ("kwonly",) + return (first_param, default_index) + except Exception: + return None + + +def _pyir_spec_open_trace(args: Any, kwargs: Any, sig: Any, entry_func: Any) -> None: + """Open the F-SPEC ledger for an outer trace: seed the re-resolvable roots + (this call's argument graph, the entry function's closure cells) and stage + the entry parameter-name roots for adoption at the entry 'fn' scope.""" + _reset_spec_state() + fn = getattr(entry_func, "__func__", entry_func) + _PYIR_SPEC_ENTRY_FUNC[0] = fn + _PYIR_SPEC_ENTRY_SIG[0] = _pyir_spec_signature_facts(fn, sig) + # Root paths key arguments by PARAMETER NAME: re-entry calls may omit + # trace-time-constant arguments, so positional indices do not re-resolve. + pos_names: "list[str]" = [] + if sig is not None: + pos_names = [ + p.name + for p in sig.parameters.values() + if p.kind + in ( + inspect.Parameter.POSITIONAL_ONLY, + inspect.Parameter.POSITIONAL_OR_KEYWORD, + ) + ] + for i, a in enumerate(args or ()): + if i >= len(pos_names): + break + name = pos_names[i] + _pyir_spec_register_root_obj(a, ("arg", name, ())) + if isinstance(a, (tuple, list)): + for j, e in enumerate(a): + _pyir_spec_register_root_obj(e, ("arg", name, (("item", j),))) + _PYIR_SPEC_PENDING_PARAM_ROOTS.append((name, ("arg", name, ()))) + if kwargs: + for name, v in kwargs.items(): + _pyir_spec_register_root_obj(v, ("arg", name, ())) + if isinstance(v, (tuple, list)): + for j, e in enumerate(v): + _pyir_spec_register_root_obj(e, ("arg", name, (("item", j),))) + _PYIR_SPEC_PENDING_PARAM_ROOTS.append((name, ("arg", name, ()))) + closure = getattr(fn, "__closure__", None) + if closure: + for ci, cell in enumerate(closure): + try: + contents = cell.cell_contents + except ValueError: + continue + _pyir_spec_register_root_obj(contents, ("closure", ci, ())) + # A scalar entry cell is a specialization input with no read + # choke of its own: record it eagerly (the callee boundary arm + # already applies exactly this contract per called function). + payload = _pyir_spec_exact_payload(contents) + if payload is not _SPEC_UNRECORDED: + _PYIR_SPEC_RECORD.setdefault(("closure", ci, ()), payload) + # Receiver candidate: a class-defined entry's first positional binding is + # its receiver (the launch binder never names it, so re-entry validation + # needs a live root). Held weakly until the seal decides it is needed. + recv = getattr(entry_func, "__self__", None) + qual_parts = getattr(fn, "__qualname__", "").split(".") + if ( + recv is None + and args + and pos_names + and len(qual_parts) >= 2 + and qual_parts[-2] != "" + ): + recv = args[0] + if recv is not None and pos_names: + try: + _PYIR_SPEC_ENTRY_RECEIVER[0] = (weakref.ref(recv), pos_names[0]) + except TypeError: + _PYIR_SPEC_ENTRY_RECEIVER[0] = None + + +def _pyir_spec_register_root_obj(obj: Any, root: tuple) -> None: + """Register *obj*'s owner token as an F-SPEC root (first-wins); staged + values and primitives carry no token and are rooted by place instead.""" + if obj is None or _is_staged_value(obj): + return + if type(obj) in (bool, int, float, str, bytes): + return + tok = _owner_token(obj) + if tok is not None: + _PYIR_SPEC_TOKEN_ROOTS.setdefault(tok, root) + + +def _pyir_spec_seg_steps(segments: "tuple[Any, ...]") -> "tuple[tuple, ...]": + """Mechanical unfold of place-key segments into re-resolution steps: a + ``_PlaceSeg`` unfolds the hops recorded at its birth, a plain str is by + definition a simple attr name (the bracket wall keeps composite strings + out of key space), anything else is an exact item key.""" + steps: "list[tuple]" = [] + for seg in segments: + if isinstance(seg, _PlaceSeg): + steps.extend(seg.steps()) + elif isinstance(seg, str): + steps.append(("attr", seg)) + else: + steps.append(("item", seg)) + return tuple(steps) + + +def _pyir_spec_root_path_for_place(place: Any) -> "tuple | None": + """Resolve a ledger place to a re-resolvable root path, or ``None`` for a + trace-internal place (a chain rooting at a trace-born object).""" + if not isinstance(place, tuple) or len(place) < 3: + return None + kind = place[0] + if kind == "local": + return _PYIR_SPEC_LOCAL_ROOTS.get(place) + if kind not in ("attr", "subscript"): + return None + tok = place[1] + if tok in _PYIR_SPEC_TRACE_BORN_TOKENS: + return None # trace-epoch owner: the place is trace-internal + root = _PYIR_SPEC_TOKEN_ROOTS.get(tok) + if root is None: + root = _pyir_spec_resolve_container_root(tok) + if root is None: + root = _pyir_spec_resolve_global_root(tok) + if root is None: + return None + if kind == "subscript": + seg = place[2] + steps = seg.steps() if isinstance(seg, _PlaceSeg) else (("item", seg),) + else: + steps = _pyir_spec_seg_steps(place[2:]) + return (root[0], root[1], root[2] + steps) + + +def _pyir_spec_resolve_container_root(tok: int) -> "tuple | None": + """Resolve a chained value token through its declared holder chain to a + re-resolvable root path (memoized on success); ``None`` for a chain that + bottoms out at a trace-internal holder.""" + steps: "list[tuple]" = [] + seen = {tok} + cur = tok + if tok in _PYIR_SPEC_TRACE_BORN_TOKENS: + return None # trace-epoch value: the chain is trace-internal + while True: + hop = _PYIR_SPEC_CONTAINER_CHAIN.get(cur) + if hop is None: + return None + cur, hop_steps = hop + if cur in _PYIR_SPEC_TRACE_BORN_TOKENS: + return None # the chain passes through a trace-epoch holder + steps = list(hop_steps) + steps + root = _PYIR_SPEC_TOKEN_ROOTS.get(cur) + if root is None and cur not in _PYIR_SPEC_CONTAINER_CHAIN: + root = _pyir_spec_resolve_global_root(cur) + if root is not None: + full = (root[0], root[1], root[2] + tuple(steps)) + _PYIR_SPEC_TOKEN_ROOTS[tok] = full + return full + if cur in seen: + return None + seen.add(cur) + + +def _pyir_spec_resolve_global_root(tok: int) -> "tuple | None": + """Lazily root an unregistered token at the entry function's module + globals by object identity; a scan miss is cached for the trace.""" + if tok in _PYIR_SPEC_UNROOTED_TOKENS: + return None + fn = _PYIR_SPEC_ENTRY_FUNC[0] + fn_globals = getattr(fn, "__globals__", None) + if fn_globals: + for name, val in list(fn_globals.items()): + if _pyir_lookup_owner_token(val) == tok: + root = ("global", name, ()) + _PYIR_SPEC_TOKEN_ROOTS[tok] = root + return root + _PYIR_SPEC_UNROOTED_TOKENS.add(tok) + return None + + +def _pyir_spec_canonical_symbol_root(base: Any) -> "tuple | None": + """The canonical re-resolvable root of a module or class object: modules + re-deref through ``sys.modules`` by their own name, classes through their + defining module and qualname (a ```` class has no canonical path). + Identity-verified against the live registry; ``None`` when unverifiable.""" + try: + if isinstance(base, types.ModuleType): + name = getattr(base, "__name__", None) + if name and _sys.modules.get(name) is base: + return ("module", name, ()) + return None + if isinstance(base, type): + mod_name = getattr(base, "__module__", None) + qual = getattr(base, "__qualname__", "") or "" + if not mod_name or not qual or "" in qual: + return None + mod = _sys.modules.get(mod_name) + if mod is None: + return None + cur: Any = mod + steps = tuple(("attr", part) for part in qual.split(".")) + for _, part in steps: + cur = getattr(cur, part, None) + if cur is base: + return ("module", mod_name, steps) + return None + except Exception: + return None + + +def _pyir_spec_root_symbol_base(base: Any, mod_name: "str | None") -> None: + """Root a symbol-read base object (first-wins): canonical module/class + paths, then the entry function's globals by identity, then the reading + module's own globals by identity (the choke declares the module).""" + tok = _owner_token(base) + if tok is None or tok in _PYIR_SPEC_TOKEN_ROOTS: + return + if tok in _PYIR_SPEC_CONTAINER_CHAIN: + return # already rooted through a holder chain + canonical = _pyir_spec_canonical_symbol_root(base) + if canonical is not None: + _PYIR_SPEC_TOKEN_ROOTS[tok] = canonical + return + if _pyir_spec_resolve_global_root(tok) is not None: + return + if not mod_name: + return + mod = _sys.modules.get(mod_name) + mod_dict = getattr(mod, "__dict__", None) + if not mod_dict: + return + for name, val in list(mod_dict.items()): + if val is base: + _PYIR_SPEC_TOKEN_ROOTS[tok] = ("module", mod_name, (("attr", name),)) + _PYIR_SPEC_UNROOTED_TOKENS.discard(tok) + return + + +def _pyir_spec_observe_attr_read( + path: Any, base: Any, attr: Any, value: Any, mod_name: "str | None" +) -> None: + """F-SPEC observation arm for an attribute read the staging chokes never + route (function-scope, test-position, and module/class-spelled reads): + roots the base symbol, chains the value hop, and records the read under + its place. Observation-only -- the read's value is never altered.""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + if base is None or _is_staged_value(base) or isinstance(base, _WatchedM): + return + _pyir_spec_root_symbol_base(base, mod_name) + _pyir_spec_chain_value(value, base, attr) + _pyir_spec_record_read(_pyir_read_place(path, base, attr), value) + except Exception: + # An unobserved bake is a staleness channel the validator cannot + # re-check: the record turns incomplete (fail closed at re-entry). + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_spec_observe_global_read( + name: Any, value: Any, mod_name: "str | None" +) -> None: + """F-SPEC observation arm for a bare global-name read (no staging choke + exists for the binding): scalar payloads record under the module root, + object payloads root the object for its downstream attr/container legs.""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + if _is_staged_value(value) or isinstance(value, _WatchedM): + return + mod = _sys.modules.get(mod_name) if mod_name else None + mod_dict = getattr(mod, "__dict__", None) + if mod_dict is None or name not in mod_dict: + return # shadowed / non-module binding: no re-resolvable fact + current = mod_dict[name] + root = ("module", mod_name, (("attr", name),)) + payload = _pyir_spec_exact_payload(value) + if payload is not _SPEC_UNRECORDED: + if current is value or (type(current) is type(value) and current == value): + _PYIR_SPEC_RECORD.setdefault(root, payload) + return + if current is not value: + return + _pyir_spec_register_root_obj(value, root) + except Exception: + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_spec_container_snapshot( + container: Any, kind: "str | None" = None +) -> "_SpecContainerSnapshot | None": + """Exact contents snapshot of a dict/list (scalar legs by value, object + legs by identity); ``None`` for a non-container or a kind mismatch.""" + + def leg(v: Any) -> Any: + p = _pyir_spec_exact_payload(v) + return p if p is not _SPEC_UNRECORDED else _SpecTraceExitObject(v) + + if isinstance(container, dict): + if kind not in (None, "dict-keys"): + return None + return _SpecContainerSnapshot( + "dict-keys", tuple(leg(k) for k in dict.keys(container)) + ) + if isinstance(container, list): + if kind not in (None, "list"): + return None + return _SpecContainerSnapshot( + "list", tuple(leg(v) for v in list.__iter__(container)) + ) + return None + + +def _pyir_spec_record_membership(container: Any) -> None: + """F-SPEC arm for a whole-container consumption (``in`` routes through the + C-level ``__contains__``, bypassing every item choke): the bake depends on + the full contents, so the container's root records a contents snapshot.""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + tok = _OWNER_TOKENS.get(container) + if tok is None: + return + root = _PYIR_SPEC_TOKEN_ROOTS.get(tok) + if root is None: + root = _pyir_spec_resolve_container_root(tok) + if root is None: + root = _pyir_spec_resolve_global_root(tok) + if root is None: + return # trace-internal container: its bake needs no re-check + snap = _pyir_spec_container_snapshot(container) + if snap is not None: + _PYIR_SPEC_RECORD.setdefault(root, snap) + except Exception: + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_spec_exact_payload(value: Any) -> Any: + """The exact structural payload a meta read bakes, or ``_SPEC_UNRECORDED`` + for a value that is not itself a baked scalar (object-shaped reads bake + through their leaf reads, which record themselves).""" + if isinstance(value, _WatchedM): + value = value.python_value + if value is None or type(value) in (bool, int, float, str, bytes): + return value + if isinstance(value, _enum.Enum): + return value + if type(value) is tuple: + elems = tuple(_pyir_spec_exact_payload(e) for e in value) + if any(e is _SPEC_UNRECORDED for e in elems): + return _SPEC_UNRECORDED + return elems + return _SPEC_UNRECORDED + + +class _SpecUnrecorded: + """Sentinel: the read's value is not a recordable exact payload.""" + + +_SPEC_UNRECORDED = _SpecUnrecorded() + + +def _pyir_spec_numeric_payload(value: Any) -> Any: + """Payload for a VALUE-SEMANTIC numeric wrapper written to a persistent + place. A wrapper carrying a constexpr scalar records that scalar BY VALUE + (its meaning is the number, not the wrapper instance -- a fresh same-value + wrapper is not stale). Returns ``None`` otherwise, so the caller keeps its + own object/identity handling (a non-derivable staged scalar and an opaque + object both stay identity rows).""" + prim = _pyir_meta_primitive_value(value) + if prim is not _NO_CONST_VALUE: + payload = _pyir_spec_exact_payload(prim) + if payload is not _SPEC_UNRECORDED: + return payload + return None + + +def _pyir_spec_record_read(place: Any, value: Any) -> None: + """F-SPEC recording funnel: a meta read at a root-pathable place records + (root path -> exact payload), first-wins. Observation-only on the trace + path; a trace-internal or non-scalar read records nothing. A value row + SUBSUMES a presence-probe row at the same path (the read bakes strictly + more than the probe's boolean).""" + try: + if place is None or not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + payload = _pyir_spec_exact_payload(value) + if payload is _SPEC_UNRECORDED: + return + root_path = _pyir_spec_root_path_for_place(place) + if root_path is None: + return + if _PYIR_SPEC_RECORD.get(root_path) is _SPEC_ATTR_PRESENT: + _PYIR_SPEC_RECORD[root_path] = payload + return + _PYIR_SPEC_RECORD.setdefault(root_path, payload) + except Exception: + # Recording must never break the read it observes, but an unrecorded + # bake is a staleness channel the validator cannot re-check: the + # record turns incomplete (fail closed at re-entry). + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_spec_record_attr_probe(owner: Any, attr: str, present: bool) -> None: + """F-SPEC arm for an attribute PROBE (``hasattr`` / 3-arg ``getattr`` + default-miss): the bake is the presence boolean, never the value, so a + root-pathable owner records a presence/absence row that re-entry re-probes + (the answer flipping post-compile refuses instead of re-entering stale).""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + place = _make_slot_key(None, owner, attr) + if place is None: + return + root_path = _pyir_spec_root_path_for_place(place) + if root_path is None: + return # trace-internal owner: its bake needs no re-check + payload = _SPEC_ATTR_PRESENT if present else _SPEC_ATTR_ABSENT + _PYIR_SPEC_RECORD.setdefault(root_path, payload) + except Exception: + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_spec_note_fabricated_bake(owner: Any, attr: str) -> None: + """F-SPEC arm for a tolerated ``__getattr__``-fabricated bare-meta read: + a ROOTED place re-derives through the hook at re-entry (the recording + funnel keeps it exact); an unrooted one is an unpathable bake the + validator cannot re-check -- the record turns incomplete (fail closed).""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + place = _make_slot_key(None, owner, attr) + if place is not None and _pyir_spec_root_path_for_place(place) is not None: + return + _PYIR_SPEC_COMPLETE[0] = False + except Exception: + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_spec_written_payload(value: Any) -> Any: + """The trace-exit payload a write amendment re-bakes: the exact structural + payload when the written value is scalar-shaped, the written object's + IDENTITY otherwise. Tuple interiors complete element-wise, so a mixed + tuple keeps exact verification where exactness exists. A raw dict/list is + the one shape whose observed binding is NOT the stored binding (container + adoption may re-point the holder slot to a watched twin after this + observation), so it stays unrecordable -- the row keeps its old payload + and refuses at re-entry, loud.""" + if isinstance(value, _WatchedM): + value = value.python_value + payload = _pyir_spec_exact_payload(value) + if payload is not _SPEC_UNRECORDED: + return payload + if type(value) is tuple: + elems = tuple(_pyir_spec_written_payload(e) for e in value) + if any(e is _SPEC_UNRECORDED for e in elems): + return _SPEC_UNRECORDED + return elems + if type(value) in (dict, list): + return _SPEC_UNRECORDED + # A constexpr-carrying numeric wrapper records BY VALUE; a non-derivable + # staged scalar and a genuine opaque object both keep an identity row. + num = _pyir_spec_numeric_payload(value) + if num is not None: + return num + return _SpecTraceExitObject(value) + + +def _pyir_spec_record_write(place: Any, value: Any) -> None: + """F-SPEC write amendment: a trace-observed write to a persistent place + (attr/subscript) re-bakes every recorded row under the written path to the + written trace-exit payload -- exact for scalars, object identity + otherwise. Re-entry then verifies the trace-EXIT state: the trace's own + pre-launch receiver mutations never read as drift, while any post-compile + change still refuses. A write never creates rows, and a row whose leaf + the written value cannot re-derive keeps its old payload (refuses at + re-entry, loud).""" + try: + if place is None or not _PYIR_SPEC_RECORD: + return + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + # A local rebind mutates the BINDING, not the persistent state a + # re-entry re-reads; only owner-place writes amend rows. + if not isinstance(place, tuple) or place[0] not in ("attr", "subscript"): + return + # A TRACE-EPOCH object written to a persistent place is trace-born: + # rows chaining through it are trace-internal (nothing at a later + # launch can change it independently of the constexpr inputs that + # built it, and re-resolving its interior would walk consumed IR). + if _pyir_spec_value_is_ir_wrapper(value): + _pyir_spec_stamp_trace_born(value) + root_path = _pyir_spec_root_path_for_place(place) + if root_path is None: + return + _pyir_spec_amend_rows_under(root_path, value) + except Exception: + pass # an unamended row refuses at re-entry -- loud, never silent + + +def _pyir_spec_record_unbind(owner: Any, attr: str) -> None: + """F-SPEC unbind amendment: a trace-observed attribute deletion re-bakes the + place's row to trace-exit ABSENCE (re-entry re-probes the final hop and + refuses when the name is back -- the same discipline as a write amendment). + Rows recorded UNDER the deleted path keep their old payloads: their storage + no longer resolves, so they refuse at re-entry, loud.""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + place = _make_slot_key(None, owner, attr) + if place is None: + return + root_path = _pyir_spec_root_path_for_place(place) + if root_path is None: + return # trace-internal owner: the deletion has no re-entry footprint + _PYIR_SPEC_RECORD[root_path] = _SPEC_ATTR_ABSENT + except Exception: + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_module_cached(name: str) -> bool: + """True when *name* is a fully initialized entry of the module cache + (LangRef 3.12 section 5.3.1): importing it again executes no module code.""" + mod = _sys.modules.get(name) + if mod is None: + return False + spec = getattr(mod, "__spec__", None) + return not getattr(spec, "_initializing", False) + + +def _pyir_import_is_cached( + module: "str | None", level: int, names: tuple, package: "str | None" +) -> bool: + """True when executing the import statement is a pure module-cache hit + (no module code runs -- the bindings are ordinary meta reads). A plain + import needs each dotted module cached; a ``from`` import needs the + resolved source cached and each name to resolve as an attribute or a + cached submodule. Unverifiable resolution answers False (the staged-CF + wall then refuses, loud).""" + try: + if module is None and level == 0: + return all(_pyir_module_cached(n) for n in names) + resolved = _pyir_spec_resolve_import(module, level, package) + if resolved is None or not _pyir_module_cached(resolved): + return False + mod = _sys.modules[resolved] + return all( + hasattr(mod, n) or _pyir_module_cached(f"{resolved}.{n}") for n in names + ) + except Exception: + return False + + +def _pyir_spec_record_module_identity(name: str) -> None: + """Root an in-body-imported module as a spec row: the compiled program owns + the trace-time import, so re-entry verifies ``sys.modules[name]`` is still + that exact module (a removed/replaced entry would re-import in Python).""" + mod = _sys.modules.get(name) + if mod is None: + _PYIR_SPEC_COMPLETE[0] = False # bound module not re-derivable: fail closed + return + _PYIR_SPEC_RECORD.setdefault(("module", name, ()), _SpecTraceExitObject(mod)) + + +def _pyir_spec_resolve_import( + module: "str | None", level: int, package: "str | None" +) -> "str | None": + """Absolute module name of an executed ``from``-import: *package* is the + defining module's declared package (a rewrite-time constant), anchoring + the relative walk exactly as the import system does (LangRef 3.12 + section 5.4.2).""" + if level == 0: + return module + if not package: + return None + import importlib.util as _importlib_util + + return _importlib_util.resolve_name("." * level + (module or ""), package) + + +def _pyir_spec_record_import( + module: "str | None", level: int, pairs: tuple, package: "str | None" +) -> None: + """F-SPEC record arm for an in-body import statement (meta flow): the + binding is a trace-time bake off the module, so each bound name records a + re-derivable spec root. A plain ``import`` (``module is None, level 0``) + binds only modules -- identity rows through ``sys.modules``. A ``from`` + import re-derives each binding off the resolved source module: scalar + payloads record value rows, object payloads root for their leaf reads; + a binding the module cannot re-derive fails closed (loud at re-entry).""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + if module is None and level == 0: + for dotted, value in pairs: + _pyir_spec_record_module_identity(dotted) + _pyir_spec_root_symbol_base(value, None) + return + resolved = _pyir_spec_resolve_import(module, level, package) + mod = _sys.modules.get(resolved) if resolved else None + if resolved is None or mod is None: + _PYIR_SPEC_COMPLETE[0] = False # unresolvable source: fail closed + return + _pyir_spec_record_module_identity(resolved) + for attr, value in pairs: + if isinstance(value, types.ModuleType): + _pyir_spec_record_module_identity(getattr(value, "__name__", resolved)) + _pyir_spec_root_symbol_base(value, None) + continue + root = ("module", resolved, (("attr", attr),)) + current = getattr(mod, attr, _SPEC_UNRECORDED) + payload = _pyir_spec_exact_payload(value) + if payload is not _SPEC_UNRECORDED: + if current is value or ( + type(current) is type(value) and current == value + ): + _PYIR_SPEC_RECORD.setdefault(root, payload) + else: + _PYIR_SPEC_COMPLETE[0] = False # binding drifted mid-trace + elif current is value: + _pyir_spec_register_root_obj(value, root) + else: + _PYIR_SPEC_COMPLETE[0] = False # binding not re-derivable + except Exception: + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_spec_amend_rows_under(root_path: tuple, value: Any) -> None: + """Re-bake every recorded row under *root_path* from *value* (the + trace-exit state of that path); a row whose leaf *value* cannot re-derive + keeps its old payload (refuses at re-entry, loud).""" + kind, key, steps = root_path + n = len(steps) + for rp in list(_PYIR_SPEC_RECORD): + if rp[0] != kind or rp[1] != key or rp[2][:n] != steps: + continue + cur = value + ok = True + for step_kind, step_key in rp[2][n:]: + try: + if step_kind == "item": + cur = cur[step_key] + elif step_kind == "attr": + # An attr step is an owner-SLOT hop: container + # adoption models dict items as owner slots, so a + # mapping at this hop resolves its slot by key. + cur = ( + cur[step_key] + if isinstance(cur, dict) + else getattr(cur, step_key) + ) + else: + ok = False + break + except Exception: + ok = False + break + if not ok: + continue + old = _PYIR_SPEC_RECORD.get(rp) + if isinstance(old, _SpecContainerSnapshot): + # A whole-container row re-bakes to the trace-exit contents (a + # shape change keeps the old snapshot: refuses at re-entry, loud). + snap = _pyir_spec_container_snapshot(cur, old.kind) + if snap is not None: + _PYIR_SPEC_RECORD[rp] = snap + continue + payload = _pyir_spec_written_payload(cur) + if payload is _SPEC_UNRECORDED: + continue + _PYIR_SPEC_RECORD[rp] = payload + + +def _pyir_spec_structural_amend(container: Any) -> None: + """F-SPEC amendment for a STRUCTURAL container mutation: legs under the + container may have moved, so re-bake its recorded rows from the container's + post-mutation contents (an orphaned row keeps its old payload and refuses + at re-entry, loud).""" + try: + if not _PYIR_SPEC_RECORD or not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + tok = _OWNER_TOKENS.get(container) + if tok is None: + return + root_path = _PYIR_SPEC_TOKEN_ROOTS.get(tok) + if root_path is None: + root_path = _pyir_spec_resolve_container_root(tok) + if root_path is None: + return + _pyir_spec_amend_rows_under(root_path, container) + except Exception: + pass # an unamended row refuses at re-entry -- loud, never silent + + +def _pyir_stamp_rewritten(fn: Any) -> Any: + """Innermost decorator emitted under a nested user ``def``: attests exactly + the raw function object compiled from this trace's instrumented AST. The + fact is identity-valued (the attested object itself), so a decorator + wrapper — even one that copies ``__dict__`` via ``functools.wraps`` — is + never attested, and an undecoratable decoration result (``@property``) + never receives a post-decoration attribute store. A generator/coroutine + def refuses instead of attesting: its body never runs at the call site, so + the stamp would certify a body the trace can never observe.""" + if isinstance(fn, types.FunctionType): + _code = getattr(fn, "__code__", None) + if inspect.isasyncgenfunction(fn) or inspect.iscoroutinefunction(fn): + raise DSLUserCodeError( + DiagId.UNSUP_ASYNC, + filename=getattr(_code, "co_filename", None), + lineno=getattr(_code, "co_firstlineno", None), + ) + if inspect.isgeneratorfunction(fn): + raise DSLUserCodeError( + DiagId.UNSUP_YIELD, + filename=getattr(_code, "co_filename", None), + lineno=getattr(_code, "co_firstlineno", None), + ) + setattr(fn, "_dsl_callee_rewritten", fn) + return fn + + +def _pyir_callee_is_rewritten(callee: Any) -> bool: + """True only for a function object a rewrite genuinely attested: the + identity-valued preprocessor stamp (bound methods unwrap to their + underlying function) or the module rewriter's ``True`` stamp on the + exec-compiled result.""" + fn = getattr(callee, "__func__", callee) + mark = getattr(fn, "_dsl_callee_rewritten", None) + return mark is fn or mark is True + + +def _pyir_spec_boundary_closure_read(func: Any, ci: int, contents: Any) -> None: + """F-SPEC boundary arm: root a callee's closure-cell meta read. A cell of + a trace-rewritten nested def is trace-internal; the entry function's own + cells root at ("closure", i); a module-global callee's cells re-deref + through the entry function's globals. Any other scalar cell is a bake + with no re-resolvable root, so the record turns incomplete (R6b).""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + payload = _pyir_spec_exact_payload(contents) + if payload is _SPEC_UNRECORDED: + return # staged / object-shaped cells are not scalar bakes + if _pyir_callee_is_rewritten(func): + return # nested def born inside this trace: cells are internal + entry = _PYIR_SPEC_ENTRY_FUNC[0] + if entry is None: + return + if func is entry: + _PYIR_SPEC_RECORD.setdefault(("closure", ci, ()), payload) + return + gname = getattr(func, "__name__", None) + entry_globals = getattr(entry, "__globals__", None) or {} + if gname is not None and entry_globals.get(gname) is func: + _PYIR_SPEC_RECORD.setdefault(("global", gname, (("cell", ci),)), payload) + return + _PYIR_SPEC_COMPLETE[0] = False + except Exception: + pass # recording must never break the read it observes + + +def _pyir_boundary_taken_defaults( + callee: Any, args: Any, kwargs: Any +) -> "list[tuple[tuple, Any]]": + """The parameter defaults this call BINDS (CPython 3.12 LangRef §8.7 / + §6.3.4: defaults evaluate once at def time; a call binds the STORED object + for every defaulted parameter the arguments leave unbound). Selectors + address ``__defaults__`` by index and ``__kwdefaults__`` by name.""" + func = getattr(callee, "__func__", callee) + code = getattr(func, "__code__", None) + if code is None: + return [] + dflts = getattr(func, "__defaults__", None) or () + kwdflts = getattr(func, "__kwdefaults__", None) or {} + if not dflts and not kwdflts: + return [] + n_pos = len(args) + (1 if getattr(callee, "__self__", None) is not None else 0) + names = code.co_varnames[: code.co_argcount] + first = code.co_argcount - len(dflts) + taken: "list[tuple[tuple, Any]]" = [] + for j, d in enumerate(dflts): + i = first + j + if i < n_pos or not (0 <= i < len(names)): + continue # bound positionally (extras spill to *args) + if i >= code.co_posonlyargcount and names[i] in kwargs: + continue # bound by keyword (a pos-only name in kwargs feeds **kw) + taken.append((("default", j), d)) + for name, d in kwdflts.items(): + if name not in kwargs: + taken.append((("kwdefault", name), d)) + return taken + + +def _pyir_spec_unwrap_steps(cand: Any, func: Any) -> "tuple | None": + """Attr steps re-deref'ing *cand*'s ``__wrapped__`` chain (the standard + functools.wraps identity) down to *func*; a direct identity match is the + empty tuple; ``None`` when the chain never reaches it.""" + steps: "tuple[tuple, ...]" = () + seen: "set[int]" = set() + cur = cand + try: + while cur is not None and id(cur) not in seen: + if cur is func: + return steps + seen.add(id(cur)) + cur = getattr(cur, "__wrapped__", None) + steps += (("attr", "__wrapped__"),) + except Exception: + return None + return None + + +def _pyir_spec_callee_defaults_root(func: Any) -> "tuple | None": + """Re-resolvable root path of a boundary callee FUNCTION object: its name + (or an identity-scanned alias) in the entry function's globals, else its + defining module + qualname walk; a decorator wrapper at the canonical + binding re-derefs through the standard ``__wrapped__`` chain -- always + identity-verified against the live object; ``None`` when no canonical + path re-derefs to this function.""" + entry = _PYIR_SPEC_ENTRY_FUNC[0] + if entry is None: + return None + entry_globals = getattr(entry, "__globals__", None) or {} + gname = getattr(func, "__name__", None) + if gname is not None and gname in entry_globals: + steps = _pyir_spec_unwrap_steps(entry_globals[gname], func) + if steps is not None: + return ("global", gname, steps) + for name, val in list(entry_globals.items()): + if val is func: + return ("global", name, ()) + if callable(val): + steps = _pyir_spec_unwrap_steps(val, func) + if steps is not None: + return ("global", name, steps) + mod_name = getattr(func, "__module__", None) + qual = getattr(func, "__qualname__", "") or "" + if mod_name and qual and "" not in qual: + cur: Any = _sys.modules.get(mod_name) + for part in qual.split("."): + cur = getattr(cur, part, None) + steps = _pyir_spec_unwrap_steps(cur, func) + if steps is not None: + return ( + "module", + mod_name, + tuple(("attr", p) for p in qual.split(".")) + steps, + ) + return None + + +def _pyir_register_taken_default_holders(callee: Any, args: Any, kwargs: Any) -> None: + """Executor-side intake for an INLINE rewritten callee: its taken defaults + enter the body as parameter bindings with no dispatcher walk in the way, + so they join the candidate-holder domain exactly like call arguments + (their leaf reads then record at the body's own chokes). Never raises.""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + for _sel, d in _pyir_boundary_taken_defaults(callee, args, kwargs): + if _is_staged_value(d): + continue + _pyir_register_candidate_holder(d) + if isinstance(d, (tuple, list)): + for e in d: + _pyir_register_candidate_holder(e) + except Exception: + pass + + +def _pyir_spec_boundary_default_reads(callee: Any, args: Any, kwargs: Any) -> None: + """F-SPEC boundary arm: a TAKEN default is external def-time state consumed + with no read choke (LangRef §8.7), so each taken payload seals at the + callee's ``__defaults__``/``__kwdefaults__`` path -- scalars by exact + payload, plain containers by contents snapshot plus per-leg rows, object + defaults by root registration for their leaf rows. A taken payload with + no re-resolvable callee path turns the record incomplete (fail closed).""" + try: + if not _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + return + taken = _pyir_boundary_taken_defaults(callee, args, kwargs) + if not taken: + return + func = getattr(callee, "__func__", callee) + if _pyir_callee_is_rewritten(func): + return # trace-born nested def: its defaults are trace-internal + root = _pyir_spec_callee_defaults_root(func) + for sel, d in taken: + if _is_staged_value(d): + continue # staged flow is not a meta bake + if root is None: + _PYIR_SPEC_COMPLETE[0] = False + continue + path = (root[0], root[1], root[2] + (sel,)) + if isinstance(d, (_WatchedDict, _WatchedList)): + # Adopted container: consumptions record at its chokes; the + # defaults path makes those rows re-resolvable. + _pyir_spec_register_root_obj(d, path) + continue + if isinstance(d, _WatchedM): + continue # place-attributed wrapper: reads record at its place + payload = _pyir_spec_exact_payload(d) + if payload is not _SPEC_UNRECORDED: + _PYIR_SPEC_RECORD.setdefault(path, payload) + continue + _pyir_spec_register_root_obj(d, path) + if isinstance(d, (dict, list)): + snap = _pyir_spec_container_snapshot(d) + if snap is not None: + _PYIR_SPEC_RECORD.setdefault(path, snap) + if isinstance(d, dict): + items: Any = [(k, dict.__getitem__(d, k)) for k in dict.keys(d)] + elif isinstance(d, tuple): + items = list(enumerate(d)) + else: + continue # instance default: leaf rows record at their chokes + for k, v in items: + leg = _pyir_spec_exact_payload(v) + _PYIR_SPEC_RECORD.setdefault( + (path[0], path[1], path[2] + (("item", k),)), + leg if leg is not _SPEC_UNRECORDED else _SpecTraceExitObject(v), + ) + except Exception: + _PYIR_SPEC_COMPLETE[0] = False + + +def _pyir_registry_candidate_objects() -> "list[Any]": + """Return the live candidate holders in sighting order (weakrefs resolved, + strong-held entries passed through).""" + out: "list[Any]" = [] + for entry in list(_PYIR_CANDIDATE_HOLDERS.values()): + obj = entry() if isinstance(entry, _weakref.ref) else entry + if obj is not None: + out.append(obj) + return out + + +def _set_slot_mv(owner: Any, slot_name: Any, mv: "MutableValue") -> "MutableValue": + """Register *mv* for the slot and return the CANONICAL cell: a live same-typed + cell already at the place is adopted (stored into), never replaced.""" + if owner is None: + if slot_name is None: + return mv + try: + place = _make_slot_key(slot_name, None, None) + except Exception: + return mv + if place is None: + return mv + _local_prev = _PLACE_REGISTRY.get(place) + if ( + _local_prev is not None + and _local_prev is not mv + and _local_prev._ref is not None + and _local_prev._is_ref_accessible() + and mv._ref is not None + and _local_prev._ref.type == mv._ref.type + ): + # Live place cell of the same type: adopt it (INV-1). + if mv._value is not None and _is_staged_value(mv._value): + _local_prev.store(mv._value) + _local_prev._place = place + return _local_prev + _PLACE_REGISTRY[place] = mv + # A fresh cell registered at the place IS the new generation: any + # superseded-unadvanced mark on the place is discharged. + _PYIR_SUPERSEDED_PLACE_ROWS.pop(place, None) + mv._place = place + return mv + # Also register the SAME MutableValue by PLACE (adopting a live same-typed + # cell); an attr place is unroll-invariant, so this runs inside constexpr + # scopes too (mirroring the read side). + _sc_place = None + try: + _place = _corrected_place_for(owner, slot_name) + if _place is not None: + _sc_place = _place + _prev = _PLACE_REGISTRY.get(_place) + if ( + _prev is not None + and _prev is not mv + and _prev._ref is not None + and _prev._is_ref_accessible() + and mv._ref is not None + and _prev._ref.type == mv._ref.type + ): + # Live place cell of the same type: adopt it (INV-1). + if mv._value is not None and _is_staged_value(mv._value): + _prev.store(mv._value) + mv = _prev + else: + _PLACE_REGISTRY[_place] = mv + _PYIR_SUPERSEDED_PLACE_ROWS.pop(_place, None) + # Stamp the cell with its place so a type-transition re-mint + # can re-book the place's bare-ref row. + mv._place = _place + except Exception: + pass + # Emission self-check post-condition: the place must now map to the ONE + # canonical cell; OUTSIDE the guard so the diagnostic raise propagates. + if _sc_place is not None: + _sc_prev = _PLACE_REGISTRY.get(_sc_place) + if _sc_prev is not None and _sc_prev is not mv: + _pyir_emission_self_check( + "write", owner, slot_name, mv, _sc_place, _sc_prev + ) + if _registry_owner(owner): + slot_id = _make_slot_id(owner, slot_name) + _SLOT_REGISTRY[slot_id] = mv + _pyir_track_slot_holder(owner) + return mv + store = _slot_store_for_tier1(owner, create=True) + if store is None: + # A dropped registration severs the slot's carry -- surface an unknown + # owner kind instead of silently skipping it. + raise DSLRuntimeError( + "PyIR internal error: tier-1 slot storage is unavailable for " + f"owner of type {type(owner).__name__!r} (slot {slot_name!r})." + ) + store[slot_name] = mv + _pyir_track_slot_holder(owner) + return mv + + +def _slot_registry_attr_mvs_for_owner( + owner: Any, +) -> "list[tuple[Any, MutableValue]]": + """Enumerate *owner*'s ATTRIBUTE slots in ``_SLOT_REGISTRY`` (registration + order); subscript slots are excluded.""" + oid = id(owner) + return [ + (slot_id.key, mv) + for slot_id, mv in list(_SLOT_REGISTRY.items()) + if slot_id.kind == "attr" and slot_id.owner == oid + ] + + +def _iter_owner_slot_mvs( + owner: Any, +) -> "list[tuple[Any, MutableValue]]": + """Yield ``(slot_name, MutableValue)`` pairs registered for *owner* (tier-1 + store or ``_SLOT_REGISTRY``); ``pyir_read`` refreshes them so property + getters see cross-region writes.""" + if owner is None: + return [] + store = _slot_store_for_tier1(owner) + if store is not None: + return list(store.items()) + if _registry_owner(owner) and not isinstance(owner, (dict, list)): + return _slot_registry_attr_mvs_for_owner(owner) + return [] + + +# Superseded-generation detection: rebind events, stamps, and read/write fire +# checks; keyed on registry facts only, never value shapes. Tables: pyir_state. + + +class _PyirGenEventSuppress: + """Reentrancy guard: while entered, generation rebind events are not + recorded and superseded-generation checks do not fire.""" + + def __enter__(self) -> None: + _PYIR_GEN_EVENT_SUPPRESS[0] += 1 + + def __exit__(self, *_exc: Any) -> "_Literal[False]": + _PYIR_GEN_EVENT_SUPPRESS[0] -= 1 + return False + + +def _pyir_gen_events_suppressed() -> bool: + return _PYIR_GEN_EVENT_SUPPRESS[0] > 0 + + +def _pyir_is_generation_compound(value: Any) -> bool: + """True for a compound user object that can carry staged leaf places -- + the object kind whose whole-object rebind is a generation event.""" + if not _has_instance_storage(value): + return False + if isinstance(value, (int, float, bool, str, bytes, type, types.ModuleType)): + return False + if _is_staged_value(value): + return False + return True + + +def _pyir_keepalive_generation_obj(obj: Any) -> None: + """Pin *obj*'s id() for the trace so the id-keyed generation tables never + alias a recycled address; weakref with drop callback when possible.""" + oid = id(obj) + if oid in _GEN_OBJ_KEEPALIVE: + return + + def _drop(_ref: Any, _oid: int = oid) -> None: + # Purge EVERY (oid, *)-keyed row this pin guards: a surviving row + # would re-anchor onto whatever object recycles the address. + _SUPERSEDED_GENERATIONS.pop(_oid, None) + _COMPOUND_BINDING_SLOTS.pop(_oid, None) + for _key in [k for k in _PYIR_CF_ATTR_FIRST_DEFS if k[0] == _oid]: + _PYIR_CF_ATTR_FIRST_DEFS.pop(_key, None) + _GEN_OBJ_KEEPALIVE.pop(_oid, None) + + try: + _GEN_OBJ_KEEPALIVE[oid] = _weakref.ref(obj, _drop) + except TypeError: + _GEN_OBJ_KEEPALIVE[oid] = obj + + +def _pyir_snapshot_generation_cells(obj: Any) -> "dict[Any, tuple]": + """Snapshot *obj*'s registered leaf cells as {slot: (MutableValue, + store_version)} -- the baseline the fire rule compares against.""" + cells: "dict[Any, tuple]" = {} + for slot_name, mv in _iter_owner_slot_mvs(obj): + if mv is None or mv._ref is None: + continue + cells[slot_name] = (mv, mv._store_version) + return cells + + +def _pyir_record_compound_binding( + target_name: Any, + owner: Any, + slot_name: Any, + new_value: Any, + filename: "str | None", + lineno: "int | None", +) -> None: + """Track compound binding places at the assign choke: a first binding records + the root place; rebinding an already-rooted compound records an alias capture.""" + if new_value is None or not _pyir_is_generation_compound(new_value): + return + if _pyir_gen_events_suppressed() or is_inside_constexpr_loop(): + return + try: + binding_slot = _make_slot_key(target_name, owner, slot_name) + except Exception: + return + if binding_slot is None: + return + oid = id(new_value) + root = _COMPOUND_BINDING_SLOTS.get(oid) + if root is None: + _COMPOUND_BINDING_SLOTS[oid] = binding_slot + _pyir_keepalive_generation_obj(new_value) + # This place now roots a fresh object: a stale capture at the same + # place no longer describes the current binding. + _ALIAS_CAPTURES.pop(binding_slot, None) + return + if root == binding_slot: + _ALIAS_CAPTURES.pop(binding_slot, None) + return + # Alias capture is only consultable for a bare-local root name (the read + # choke discriminates on the dotted target's root NAME). + if binding_slot[0] != "local": + return + _ALIAS_CAPTURES[binding_slot] = { + "root": root, + "gen": _BINDING_GEN.get(root, 0), + } + _ALIAS_CAPTURE_ROOT_NAMES.add(binding_slot[2]) + + +def _pyir_stamp_superseded_generation( + obj: Any, + live_obj: Any, + target_name: Any, + filename: "str | None", + lineno: "int | None", + cells: "dict[Any, tuple]", +) -> None: + """Stamp *obj* as a superseded generation: record the pre-rebind cell + versions and its own staged leaf wrappers (excluding shared ones).""" + rec: "dict[str, Any]" = { + "target": target_name, + "site": (filename, lineno), + "cells": cells, + "wrapper_refs": [], + } + live_d = _instance_storage_items(live_obj) or {} + for slot_name, value in list((_instance_storage_items(obj) or {}).items()): + if slot_name not in cells: + continue + if not _is_staged_value(value): + continue + if live_d.get(slot_name) is value: + continue # wrapper shared with the live generation + _SUPERSEDED_LEAF_WRAPPERS[id(value)] = (rec, slot_name) + rec["wrapper_refs"].append(value) + _SUPERSEDED_GENERATIONS[id(obj)] = rec + _pyir_keepalive_generation_obj(obj) + + +# Store-version sentinel for a cell born at the rebind itself: any read +# resolving it diverges, so its baseline is always "advanced". +_PYIR_GEN_CELL_BORN_AT_REBIND = -1 + + +def _pyir_record_rebind_ref_cell( + rebind_rec: "dict | None", attr: Any, ref: Any +) -> None: + """Register a D1-promoted meta-field ``pyir.ref`` as a born-at-rebind cell + of a whole-object rebind record (the raw-ref cell form the alias-capture + channel resolves against the live meta-value-table row).""" + if rebind_rec is None or ref is None: + return + rebind_rec["cells"][attr] = (ref, _PYIR_GEN_CELL_BORN_AT_REBIND) + + +def _pyir_record_generation_rebind( + binding_slot: Any, + old_value: Any, + new_value: Any, + target_name: Any, + filename: "str | None", + lineno: "int | None", + superseded_obj: Any = None, + allow_empty_cells: bool = False, +) -> "dict | None": + """Record a whole-object rebind of a compound binding place: bump the place's + generation counter(s) and stamp a retained non-current generation.""" + if _pyir_gen_events_suppressed() or is_inside_constexpr_loop(): + return None + cells = _pyir_snapshot_generation_cells(old_value) + if not cells and not allow_empty_cells: + return None + keys = [] + if binding_slot is not None: + keys.append(binding_slot) + root = _COMPOUND_BINDING_SLOTS.get(id(old_value)) + if root is not None and root != binding_slot: + keys.append(root) + if not keys: + return None + rebind_rec = {"cells": cells, "site": (filename, lineno)} + for k in keys: + _BINDING_GEN[k] = _BINDING_GEN.get(k, 0) + 1 + _BINDING_REBIND_CELLS[k] = rebind_rec + if superseded_obj is not None and cells: + _pyir_stamp_superseded_generation( + superseded_obj, new_value, target_name, filename, lineno, cells + ) + return rebind_rec + + +def _pyir_complete_generation_rebind_cells( + rebind_rec: "dict | None", old_value: Any +) -> None: + """Complete an m2m rebind record with cells the decompose walk minted; + their baseline is the always-advanced born-at-rebind sentinel.""" + if rebind_rec is None or _pyir_gen_events_suppressed(): + return + cells = rebind_rec["cells"] + for slot_name, mv in _iter_owner_slot_mvs(old_value): + if slot_name in cells or mv is None or mv._ref is None: + continue + cells[slot_name] = (mv, _PYIR_GEN_CELL_BORN_AT_REBIND) + + +def _pyir_raise_superseded(access: str, var: Any, rebind_site: "tuple | None") -> None: + """Raise the superseded-generation diagnostic.""" + detail = "" + if rebind_site and rebind_site[0]: + detail = f" (replaced at {rebind_site[0]}:{rebind_site[1]})" + raise DSLUserCodeError( + DiagId.SCOPE_READ_OF_SUPERSEDED_GENERATION, + var=str(var), + access=access, + detail=detail, + ) + + +def _pyir_check_superseded_owner_write( + owner: Any, slot_name: Any, new_value: Any, target_name: Any +) -> None: + """Write choke: a write through a superseded generation lands in the one + live cell and would corrupt the current generation -> diagnostic.""" + rec = _SUPERSEDED_GENERATIONS.get(id(owner)) + if rec is None or _pyir_gen_events_suppressed(): + return + if slot_name not in rec["cells"] and not _is_staged_value(new_value): + return + _pyir_raise_superseded("written", target_name or rec["target"], rec["site"]) + + +def _pyir_check_superseded_wrapper_load(value: Any) -> None: + """Auto-load choke: a superseded generation's own leaf wrapper whose shared + cell advanced past the stamp would observe the replacing value -> raise.""" + if not _SUPERSEDED_LEAF_WRAPPERS or value is None: + return + ent = _SUPERSEDED_LEAF_WRAPPERS.get(id(value)) + if ent is None or _pyir_gen_events_suppressed(): + return + rec, w_slot = ent + cell = rec["cells"].get(w_slot) + if cell is None: + return + mv0, v0 = cell + if mv0._store_version <= v0: + return + _pyir_raise_superseded("read", rec["target"], rec["site"]) + + +def _pyir_check_alias_generation_read( + root_name: str, target_name: Any, owner: Any, slot_name: Any +) -> None: + """Alias-capture channel: fires only when the dotted read resolves the same + cell the root place held at the rebind AND its store version advanced.""" + try: + key = _make_slot_key(root_name, None, None) + except Exception: + return + cap = _ALIAS_CAPTURES.get(key) + if cap is None: + return + root_slot = cap["root"] + if _BINDING_GEN.get(root_slot, 0) <= cap["gen"]: + return + rebind = _BINDING_REBIND_CELLS.get(root_slot) + if rebind is None: + return + cell = rebind["cells"].get(slot_name) + if cell is None: + return + mv0, v0 = cell + if isinstance(mv0, ir.Value): + # Raw-ref cell: a D1-promoted meta field of a carried all-meta + # compound rebind. Identity is the live meta-value-table row; the + # born-at-rebind baseline means any capture-crossing read diverges. + try: + live = _slot_refs.get(_make_slot_key(None, owner, slot_name)) + except Exception: + live = None + if not isinstance(live, ir.Value) or not _same_ir_value(live, mv0): + return + _pyir_raise_superseded("read", target_name, rebind["site"]) + mv = _get_slot_mv(owner, slot_name) + if mv is not mv0: + return + if mv0._store_version <= v0: + return + _pyir_raise_superseded("read", target_name, rebind["site"]) + + +def _pyir_generation_read_checks( + target_name: Any, current_value: Any, owner: Any, slot_name: Any +) -> None: + """Read choke: owner-identity, leaf-wrapper, and alias-capture channels, + each guarded by the store-version-advanced conjunct.""" + if _pyir_gen_events_suppressed(): + return + if owner is not None and _SUPERSEDED_GENERATIONS: + rec = _SUPERSEDED_GENERATIONS.get(id(owner)) + if rec is not None: + cell = rec["cells"].get(slot_name) + if cell is not None: + mv0, v0 = cell + if mv0._store_version > v0: + _pyir_raise_superseded( + "read", target_name or rec["target"], rec["site"] + ) + _pyir_check_superseded_wrapper_load(current_value) + if _ALIAS_CAPTURE_ROOT_NAMES and owner is not None and isinstance(target_name, str): + root_name, sep, _rest = target_name.partition(".") + if sep and root_name in _ALIAS_CAPTURE_ROOT_NAMES: + _pyir_check_alias_generation_read(root_name, target_name, owner, slot_name) + + +# Region-conditional attribute first-defs: the attr persists on the shared +# owner beyond its arm, so an uncovered read re-materializes an unbound value. + + +def _pyir_current_region_arm_path() -> tuple: + """The trace-time region-arm path: the 'region' frames of the scope stack; + distinct frame objects distinguish sibling arms of one ``scf.if``.""" + return tuple(f for f in _PYIR_SCOPE_STACK if f.kind == "region") + + +def _pyir_arm_path_covers(recorded: tuple, current: tuple) -> bool: + """True iff every arm on the recorded path is still open on the current + path (identity-compared prefix).""" + if len(recorded) > len(current): + return False + for a, b in zip(recorded, current): + if a is not b: + return False + return True + + +def _pyir_record_cf_attr_first_def(owner: Any, slot_name: Any) -> None: + """Record a staged attr-leaf first-def minted inside a staged region; keys + on binding-position facts only (owner identity + region-arm path).""" + if owner is None or slot_name is None: + return + if _pyir_gen_events_suppressed() or is_inside_constexpr_loop(): + return + arms = _pyir_current_region_arm_path() + if not arms: + return # not inside a region body: unconditional first-def + src_file, src_line = _first_non_dsl_caller_location() + _PYIR_CF_ATTR_FIRST_DEFS[(id(owner), slot_name)] = { + "arms": arms, + "site": (src_file, src_line), + } + _pyir_keepalive_generation_obj(owner) + # Under a folded constant if branch, hand the record to the latch consumer: + # a latched gate proves a once-per-key init and discharges it. + if _PYIR_FOLD_FIRSTDEF_STACK: + _PYIR_FOLD_FIRSTDEF_STACK[-1].append( + (CF_ATTR_FIRST_DEF, (id(owner), slot_name)) + ) + + +def _pyir_update_cf_attr_first_def_on_write(owner: Any, slot_name: Any) -> None: + """A later write outside staged CF drops the record; a write outside the + recorded subtree re-anchors it at the wider arm path.""" + if owner is None or slot_name is None or not _PYIR_CF_ATTR_FIRST_DEFS: + return + key = (id(owner), slot_name) + rec = _PYIR_CF_ATTR_FIRST_DEFS.get(key) + if rec is None or _pyir_gen_events_suppressed(): + return + if not is_inside_staged_cf(): + _PYIR_CF_ATTR_FIRST_DEFS.pop(key, None) + return + cur_arms = _pyir_current_region_arm_path() + if not cur_arms: + _PYIR_CF_ATTR_FIRST_DEFS.pop(key, None) + return + if not _pyir_arm_path_covers(rec["arms"], cur_arms): + rec["arms"] = cur_arms + + +def _pyir_check_cf_attr_first_def_read(owner: Any, slot_name: Any) -> None: + """Read choke: raise when the read site is outside the first-def's arm-path + subtree (sibling arm or post-region), where Python never bound the attr.""" + if owner is None or slot_name is None or not _PYIR_CF_ATTR_FIRST_DEFS: + return + rec = _PYIR_CF_ATTR_FIRST_DEFS.get((id(owner), slot_name)) + if rec is None or _pyir_gen_events_suppressed(): + return + if is_inside_constexpr_loop(): + return + if _pyir_arm_path_covers(rec["arms"], _pyir_current_region_arm_path()): + return + detail = "" + def_file, def_line = rec["site"] + if def_file: + detail = f" (only set at {def_file}:{def_line}, inside a branch that does not cover this read)" + raise DSLUserCodeError(DiagId.SCOPE_READ_NEVER_SET, detail=detail) + + +def _attach_mutable_ref(obj: object, mv: "MutableValue", context: str) -> None: + """Attach *mv* as ``_mutable_ref`` on *obj*; logs (never raises) when the + object does not accept arbitrary attributes.""" + try: + _pyir_setattr_raw(obj, "_mutable_ref", mv) + except (AttributeError, TypeError): + log().info( + "could not attach _mutable_ref to %s (%s)", + type(obj).__name__, + context, + ) + + +def _pyir_adopt_live_place_cell(owner: Any, slot_name: Any, value: Any) -> Any: + """Bind *value* at a live place cell: store into it, re-alias ``_mutable_ref`` + and the id tier to it. Never mints cells; steps aside on a type change.""" + if owner is None or slot_name is None or value is None: + return value + try: + if not (_is_staged_value(value) and _can_carry_leaf_ref(value)): + return value + place = _corrected_place_for(owner, slot_name) + if place is None: + return value + place_mv = _PLACE_REGISTRY.get(place) + if place_mv is None or place_mv._ref is None: + return value + if not place_mv._is_ref_accessible(): + return value + if getattr(value, "_mutable_ref", None) is place_mv: + return value + if _pyir_ref_pointee_type_changed(place_mv._ref, value): + return value + # A backing trapped in an already-closed region cannot be stored at the + # current IP (conservative skip). + if _raw_backing_ir_value(value) is not None and not _value_dominates_current_ip( + value + ): + return value + place_mv.store(value) + # Clear any foreign-cell load tag so the next auto-load reads the place + # cell, not the foreign cell's SSA. + try: + _pyir_setattr_raw(value, _PYIR_LOAD_VERSION_ATTR, None) + except (AttributeError, TypeError): + pass + _attach_mutable_ref(value, place_mv, f"place-cell binding '{slot_name}'") + _set_slot_mv(owner, slot_name, place_mv) + except Exception: + return value + return value + + +def _pyir_place_is_index_sibling(route_place: Any, place: Any) -> bool: + """True when the two ledger places name elements of the SAME parent path + and differ only in the final integer index -- the positional-shift + signature of a shrinking/growing restructure. A whole-value place vs an + element place ('t' vs 't[0]') is NOT a sibling pair (LAW-2: re-aliasing + across it would steal another live binding's route).""" + if ( + not isinstance(route_place, tuple) + or not isinstance(place, tuple) + or len(route_place) != len(place) + or route_place == place + ): + return False + diff = [i for i in range(len(place)) if route_place[i] != place[i]] + if len(diff) != 1: + return False + a, b = route_place[diff[0]], place[diff[0]] + if isinstance(a, _PlaceSeg) and isinstance(b, _PlaceSeg): + return ( + a.base == b.base + and len(a.keys) == len(b.keys) + and bool(a.keys) + and a.keys[:-1] == b.keys[:-1] + and isinstance(a.keys[-1], int) + and isinstance(b.keys[-1], int) + ) + if isinstance(a, str) and isinstance(b, str): + base_a, sep_a, idx_a = a.rpartition("[") + base_b, sep_b, idx_b = b.rpartition("[") + return ( + bool(sep_a) + and bool(sep_b) + and base_a == base_b + and idx_a.endswith("]") + and idx_b.endswith("]") + and idx_a[:-1].isdigit() + and idx_b[:-1].isdigit() + ) + return False + + +def _pyir_route_restructured_tuple_leaves( + target_name: Any, owner: Any, slot_name: Any, new_value: Any +) -> None: + """Generation coherence for a restructuring rebind: each element of the + fresh tuple stores through its access path's still-live cell, so a + post-join read of ``name[i]`` resolves the fresh generation. A live cell + the write cannot advance is marked superseded-unadvanced; serving it to a + later place-routed read refuses instead of reading the old generation.""" + for i, elem in enumerate(new_value): + elem_name = f"{target_name}[{i}]" + elem_slot = _place_seg_child(slot_name, i) if slot_name is not None else None + if isinstance(elem, (tuple, list)): + _pyir_route_restructured_tuple_leaves( + elem_name, owner, elem_slot if owner is not None else None, elem + ) + continue + try: + if owner is not None and elem_slot is not None: + mv = _get_slot_mv(owner, elem_slot) + place = _make_slot_key(None, owner, elem_slot) + else: + mv = _get_slot_mv(None, elem_name) + place = _make_slot_key(elem_name, None, None) + except Exception: + continue + if mv is None or mv._ref is None or not mv._is_ref_accessible(): + log().info( + "[pyir_assign] '%s' restructure write-through: no live row", + elem_name, + ) + continue + can_store = ( + _is_staged_value(elem) + and _can_carry_leaf_ref(elem) + # LAW-1 position rule: a literal-backed value re-materializes at + # the store position, so only SSA-backed values need dominance. + and (_is_literal_backed(elem) or _value_dominates_current_ip(elem)) + ) + if can_store: + if _is_literal_backed(elem): + # Type-probe on a rebuilt twin: ``ir_value()`` on the element + # itself would demote it from literal- to SSA-backed. + try: + probe_raw = type(elem)(elem.value).ir_value() + can_store = probe_raw.type == mv._ref.type.pointee + except Exception: + can_store = False + else: + can_store = not _pyir_ref_pointee_type_changed(mv._ref, elem) + if can_store: + mv.store(elem) + if place is not None: + _PYIR_SUPERSEDED_PLACE_ROWS.pop(place, None) + # An index-shifted survivor still routes to its OLD sibling row, + # which this same pass may overwrite with a higher index's leaf; + # re-alias it to the row it now lives at (mirroring + # _pyir_adopt_live_place_cell's binding step). Only the + # index-sibling shift re-aliases: a route to any other place is + # another live binding's (LAW-2) and keeps it. + route = getattr(elem, "_mutable_ref", None) + if ( + route is not None + and route is not mv + and _pyir_place_is_index_sibling(getattr(route, "_place", None), place) + ): + try: + _pyir_setattr_raw(elem, _PYIR_LOAD_VERSION_ATTR, None) + except (AttributeError, TypeError): + pass + _attach_mutable_ref(elem, mv, f"restructure shift '{elem_name}'") + log().info( + "[pyir_assign] '%s' restructure write-through: stored", elem_name + ) + elif place is not None: + _PYIR_SUPERSEDED_PLACE_ROWS[place] = mv._store_version + log().info( + "[pyir_assign] '%s' restructure write-through: cell not " + "advanceable → superseded mark", + elem_name, + ) + + +def _pyir_refuse_superseded_row_serve(mv: Any, target_name: Any) -> None: + """The reroute guard: a place-routed read about to be served by a row whose + cell a restructuring rebind could NOT advance would read the superseded + generation -- refuse loudly. Any tracked store since the mark (a differing + store version) is the new generation's write and discharges the mark.""" + if not _PYIR_SUPERSEDED_PLACE_ROWS: + return + place = getattr(mv, "_place", None) + if place is None: + return + marked_version = _PYIR_SUPERSEDED_PLACE_ROWS.get(place) + if marked_version is None: + return + if mv._store_version != marked_version: + _PYIR_SUPERSEDED_PLACE_ROWS.pop(place, None) + return + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.UNOBSERVED_WRITE_POSITION_UNKNOWN, + filename=filename, + lineno=lineno, + var=str(target_name), + ) + + +def _fresh_wrapper(dsl_value: Any) -> Any: + """Return a fresh DSL wrapper sharing *dsl_value*'s backing (distinct object for + the wrapper-keyed ref cache, same ``.value``); unchanged on any failure.""" + # Literal-backed Numeric: rebuild from the raw Python scalar so the wrapper stays literal-backed; + # ``ir_value()`` would emit an eager ``arith.constant``, demoting it to SSA-backed. + if _is_literal_backed(dsl_value): + try: + return type(dsl_value)(dsl_value.value) + except Exception: + return dsl_value + + try: + ir_val = dsl_value.ir_value() + except Exception: + return dsl_value + + # Rebuild through the ONE reconstruction funnel (the copy is a best-effort + # convenience, not a read: an unrebuildable value passes through unchanged). + try: + fresh = _wrap_ir_like(dsl_value, ir_val) + except DSLUserCodeError: + return dsl_value + + # Re-point the rebuilt wrapper's ``.value`` at the original ``ir.Value`` so a + # no-op cast keeps ``b.value is a.value``; SSA-backed scalars only. + src_value = getattr(dsl_value, "value", None) + if ( + fresh is not dsl_value + and isinstance(src_value, ir.Value) + and isinstance(getattr(fresh, "value", None), ir.Value) + ): + try: + _pyir_setattr_raw(fresh, "value", src_value) + except (AttributeError, TypeError): + pass + # Mark a local-owned wrapper so ``pyir_assign`` can preserve Python ``is`` + # identity instead of re-wrapping a compound handed back by a no-op op. + if fresh is not dsl_value: + try: + _pyir_setattr_raw(fresh, "_pyir_local_wrapper", True) + except (AttributeError, TypeError): + pass + return fresh + + +def _same_ir_value(a: "ir.Value | None", b: "ir.Value | None") -> bool: + """Whether *a* and *b* are the same SSA value, via the base ``ir.Value.__eq__`` + (a subclass ``==`` may emit IR). Conservative: any error -> ``False``.""" + if a is None or b is None: + return False + try: + return bool(ir.Value.__eq__(a, b)) + except Exception: + return a is b + + +def _raw_backing_ir_value(arg: Any) -> "ir.Value | None": + """Return *arg*'s already-baked backing ``ir.Value`` structurally, never via + ``ir_value()`` (which re-enters the auto-load path); ``None`` when unavailable.""" + if isinstance(arg, ir.Value): + return arg + inner = getattr(arg, "value", None) + if isinstance(inner, ir.Value): + return inner + return None + + +def _as_arith_capable_scalar_leaf(value: Any) -> Any: + """Wrap a bare SCALAR int/float ``ir.Value`` in ``ArithValue`` so downstream + staged arithmetic works; opaque leaves and non-``ir.Value``s pass through.""" + # Local import: ``base_dsl.typing`` imports ``_WatchedM`` from this module, so a + # module-level import of ``ArithValue`` would be circular. + from .._mlir_helpers.arith import ArithValue + + if not isinstance(value, ir.Value) or isinstance(value, ArithValue): + return value + try: + ty = value.type + if isinstance(ty, (ir.IntegerType, ir.FloatType)): + return ArithValue(value) + except Exception: + pass + return value + + +def _value_dominates_current_ip(value: Any) -> bool: + """True if *value*'s backing ``ir.Value`` dominates the current insertion point + (C++ binding). Conservative: ``False`` on any error, so callers re-load.""" + if pyir is None: + return False + try: + raw = _raw_backing_ir_value(value) + if raw is None: + return False + ip = ir.InsertionPoint.current + ref_op = ip.ref_operation + if ref_op is not None: + ref_op = getattr(ref_op, "operation", ref_op) + return bool(pyir.value_dominates_ip(raw, ip.block, ref_op)) + except Exception: + return False + + +def _ref_literal_init_const(ref: "ir.Value") -> Any: + """Return the literal a ``pyir.ref`` was GENUINELY seeded from, or + ``_NO_CONST_VALUE`` (a placeholder/poison init is not a seed). C++-backed.""" + try: + if pyir is None: + return _NO_CONST_VALUE + result = pyir.ref_literal_init(ref) + except Exception: + return _NO_CONST_VALUE + return _NO_CONST_VALUE if result is None else result + + +def _op_is_inside_op(op: "ir.Operation", outer_op: "ir.Operation") -> bool: + """True if *op* is nested (at any depth) inside *outer_op*, via the + ``Operation.parent`` chain only. Conservative: any error -> ``False``.""" + try: + target = getattr(outer_op, "operation", outer_op) + cur = getattr(op, "operation", op) + while cur is not None: + if cur == target: + return True + parent = cur.parent + if parent is None: + return False + cur = getattr(parent, "operation", parent) + return False + except Exception: + return False + + +def _block_strictly_inside(inner: "ir.Block", outer: "ir.Block") -> bool: + """True if *inner* block is nested strictly inside *outer* block (owning-op + containment, blocks differ). Conservative: False.""" + try: + if inner is None or outer is None or inner == outer: + return False + inner_op = inner.owner + outer_op = outer.owner + if inner_op is None or outer_op is None: + return False + return _op_is_inside_op( + getattr(inner_op, "operation", inner_op), + getattr(outer_op, "operation", outer_op), + ) + except Exception: + return False + + +def _block_inside_op(block: "ir.Block", op: Any) -> bool: + """True if *block* is one of *op*'s own blocks or nested inside them. + Conservative: False.""" + try: + if block is None or op is None: + return False + owner = block.owner + if owner is None: + return False + owner_op = getattr(owner, "operation", owner) + target_op = getattr(op, "operation", op) + if owner_op == target_op: + return True + return _op_is_inside_op(owner_op, target_op) + except Exception: + return False + + +def _pyir_write_in_if_arms_of(inner: "ir.Block", outer: "ir.Block") -> bool: + """True when walking up from *inner* to *outer* crosses ONLY ``scf.if`` + regions: the write sits in conditional arms of the slot's own birth + region, with no loop in between to multiply it. Conservative: False.""" + try: + if inner is None or outer is None or inner == outer: + return False + target = outer.owner + target = getattr(target, "operation", target) + cur = inner.owner + cur = getattr(cur, "operation", cur) + for _ in range(64): + if cur is None or target is None: + return False + if cur == target: + return True + if cur.name != "scf.if": + return False + parent = cur.parent + if parent is None: + return False + cur = getattr(parent, "operation", parent) + return False + except Exception: + return False + + +def _pyir_check_arm_local_escape(slot_key: Any) -> None: + """Refuse a consumption of an arm-locally mutated meta slot from OUTSIDE + its arm: the write ran on one traced path only, so the folded value is + path-dependent at run time.""" + entry = _PYIR_ARM_LOCAL_META_WRITES.get(slot_key) + if entry is None: + return + arm_block, w_file, w_line = entry + try: + cur = ir.InsertionPoint.current.block + except Exception: + return + if cur == arm_block or _block_strictly_inside(cur, arm_block): + return + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.PHASE_ARM_LOCAL_CONSTANT_ESCAPES, + filename=filename, + lineno=lineno, + var=str(slot_key[-1]) if isinstance(slot_key, tuple) else str(slot_key), + write_file=w_file, + write_line=w_line, + ) + + +def _ir_value_defined_inside_op(value: "ir.Value", outer_op: "ir.Operation") -> bool: + """True if *value* is defined (at any depth) inside *outer_op*. + Conservative: False (treats the value as outside).""" + try: + owner = value.owner + if isinstance(owner, ir.Block): + parent = owner.owner # op owning the block + if parent is None: + return False + return _op_is_inside_op(getattr(parent, "operation", parent), outer_op) + def_op = getattr(owner, "operation", owner) + return _op_is_inside_op(def_op, outer_op) + except Exception: + return False + + +def _op_has_enclosing_loop(op: "ir.Operation") -> bool: + """True if *op* has an enclosing ``scf.for``/``scf.while`` ancestor (gates how + the ``scf.if`` forward-carry treats mutated leaves). Conservative: False.""" + try: + cur = getattr(op, "operation", op) + # Skip the op itself; only ancestors count as "enclosing". + parent = cur.parent + while parent is not None: + pop = getattr(parent, "operation", parent) + if getattr(pop, "name", None) in ("scf.for", "scf.while"): + return True + parent = pop.parent + return False + except Exception: + return False + + +def _innermost_enclosing_loop_op_at_ip() -> "ir.Operation | None": + """Innermost ``scf.for``/``scf.while`` op enclosing the current insertion + point, or ``None``; bounded upward walk stopping at function/module boundaries.""" + try: + block = ir.InsertionPoint.current.block + except Exception: + return None + + _MAX_NESTING = 256 + for _ in range(_MAX_NESTING): + if block is None: + return None + try: + parent = block.owner # Block -> Python dialect op + except Exception: + return None + op = getattr(parent, "operation", parent) # -> ir.Operation + try: + op_name = str(op.name) + except Exception: + return None + if op_name in ("scf.for", "scf.while"): + return op + # A function / module boundary owns no enclosing loop; stop rather than + # walk into recycled top-of-region block aliasing (a SIGSEGV hazard). + if _is_func_boundary_op(op_name) or _is_module_boundary_op(op_name): + return None + try: + block = op.block # ir.Operation -> parent Block + except Exception: + return None + return None + + +def _loop_free_enclosing_if_ops_at_ip() -> "list[ir.Operation]": + """Chain of loop-free ``scf.if`` ancestors of the current insertion point, + innermost outward, stopping at the first loop or function/module boundary.""" + ifs: "list[ir.Operation]" = [] + try: + block = ir.InsertionPoint.current.block + except Exception: + return ifs + + _MAX_NESTING = 256 + for _ in range(_MAX_NESTING): + if block is None: + return ifs + try: + parent = block.owner + except Exception: + return ifs + op = getattr(parent, "operation", parent) + try: + op_name = str(op.name) + except Exception: + return ifs + if op_name in ("scf.for", "scf.while"): + return ifs + if op_name == "scf.if": + ifs.append(op) + if _is_func_boundary_op(op_name) or _is_module_boundary_op(op_name): + return ifs + try: + block = op.block + except Exception: + return ifs + return ifs + + +def _region_block_of_if_containing_op( + start_op: "ir.Operation", if_op: "ir.Operation" +) -> "ir.Block | None": + """The *if_op* region block (then/else) that transitively contains *start_op*, + via the ``Operation.parent`` walk. Conservative: ``None``.""" + try: + target = getattr(if_op, "operation", if_op) + cur = getattr(start_op, "operation", start_op) + while cur is not None: + parent = cur.parent + if parent is None: + return None + parent = getattr(parent, "operation", parent) + if parent == target: + return cur.block + cur = parent + return None + except Exception: + return None + + +def _meta_use_in_sibling_if_region( + baked_uses: "list[ir.Value]", if_op: "ir.Operation" +) -> bool: + """True if a use in *baked_uses* lives in a DIFFERENT branch of *if_op* than the + current IP (sibling-branch re-trace, not a carry); any error -> ``False``.""" + if pyir is None: + return False + try: + cur_block = ir.InsertionPoint.current.block + except Exception: + return False + if cur_block is None: + return False + if_op = getattr(if_op, "operation", if_op) + # Reduce the current IP to its own branch block under ``if_op``: an op directly in ``cur_block`` + # reduces to ``cur_block`` itself; otherwise climb to the region block whose parent op is ``if_op``. + cur_branch = None + try: + cb_owner = getattr(cur_block.owner, "operation", cur_block.owner) + if cb_owner == if_op: + cur_branch = cur_block + else: + cur_branch = _region_block_of_if_containing_op(cb_owner, if_op) + except Exception: + cur_branch = None + if cur_branch is None: + return False + for use in baked_uses: + try: + if isinstance(getattr(use, "owner", None), ir.Block): + continue # block argument: no defining op + def_op = getattr(use.owner, "operation", use.owner) + use_branch = _region_block_of_if_containing_op(def_op, if_op) + if use_branch is None: + continue # baked outside this if -> not a sibling rebake + if use_branch != cur_branch: + return True + except Exception: + continue + return False + + +def _safe_instance_dict(obj: Any) -> "dict | None": + """*obj*'s instance ``__dict__`` or ``None``, read via + ``object.__getattribute__`` so an overloaded attribute access cannot abort the caller.""" + try: + d = object.__getattribute__(obj, "__dict__") + except Exception: + return None + return d if isinstance(d, dict) else None + + +# Per-class cache of declared ``__slots__`` member descriptors. Classes are +# module-lifetime objects, so a plain dict keyed by the class is safe. +_SLOTS_MEMBERS_CACHE: "dict[type, tuple[tuple[str, Any], ...]]" = {} + + +def _slots_member_descriptors(cls: type) -> "tuple[tuple[str, Any], ...]": + """(name, descriptor) for every ``__slots__`` entry across *cls*'s MRO (cached); + pseudo-slots, dunders and shadowed names are skipped.""" + cached = _SLOTS_MEMBERS_CACHE.get(cls) + if cached is not None: + return cached + members: "list[tuple[str, Any]]" = [] + seen: "set[str]" = set() + for klass in getattr(cls, "__mro__", ()): + raw = klass.__dict__.get("__slots__") + if raw is None: + continue + if isinstance(raw, str): + raw = (raw,) + try: + names = list(raw) + except TypeError: + continue + for name in names: + if not isinstance(name, str) or name.startswith("__") or name in seen: + continue + desc = klass.__dict__.get(name) + if desc is None or not hasattr(desc, "__get__"): + continue + seen.add(name) + members.append((name, desc)) + result = tuple(members) + _SLOTS_MEMBERS_CACHE[cls] = result + return result + + +def _has_instance_storage(obj: Any) -> bool: + """True when *obj* has per-instance attribute storage: a readable instance + ``__dict__`` OR declared ``__slots__`` members.""" + if _safe_instance_dict(obj) is not None: + return True + return bool(_slots_member_descriptors(type(obj))) + + +def _instance_storage_items(obj: Any) -> "dict[str, Any] | None": + """``{name: value}`` of *obj*'s instance storage (``__dict__`` entries plus + SET ``__slots__`` fields, unset slots omitted), or ``None`` when it has none.""" + d = _safe_instance_dict(obj) + members = _slots_member_descriptors(type(obj)) + if not members: + return d + items: "dict[str, Any]" = dict(d) if d else {} + for name, desc in members: + if name in items: + continue + try: + items[name] = desc.__get__(obj, type(obj)) + except Exception: + continue # unset slot (AttributeError) / exotic descriptor + return items + + +# MutableValue -- internal bookkeeping for pyir.ref / pyir.load / pyir.store + + +class MutableValue: + """Wraps a single leaf DSL value and holds the ``pyir.ref`` handle; never + exposed to user code and never participates in operator dispatch.""" + + __slots__ = ( + "_value", + "_type", + "_ref", + "_ref_context_id", + "_load_version", + "_store_version", + "_place", + "_last_choke_region_epoch", + ) + + def __init__(self, value: Any) -> None: + if isinstance(value, (bool, int, float)): + raise DSLRuntimeError( + f"Cannot create a mutable reference for Python scalar " + f"`{value}` (type: {type(value).__name__}).", + suggestion=( + "Convert to a DSL type first, e.g. " + "cutlass.Int32(...) or cutlass.Float32(...)." + ), + ) + + self._value = value + self._type = type(value) + self._ref = None # populated by take_reference() + self._ref_context_id: int | None = None + # Bumped on every load()/store(); used to dedup redundant auto-loads. A value + # whose ``_pyir_load_version`` equals the current one is the freshest load. + self._load_version: int = 0 + # Bumped ONLY on store(); the superseded-generation detector fires only + # when the shared cell demonstrably advanced past a recorded stamp. + self._store_version: int = 0 + # The ledger place this cell is registered under (``None`` if + # unregistered); a type-transition re-mint re-books the place through it. + self._place: Any = None + # Region-epoch at this cell's last store()/load() choke; equality with + # the current epoch proves a straight-line interval since that event. + self._last_choke_region_epoch: "int | None" = None + + def take_reference(self, *, remint: bool = False) -> None: + """Create the ``pyir.ref`` for the current value -- exactly once per trace; + re-minting an accessible cell is an internal error unless *remint*.""" + if self._ref is not None and not remint and self._is_ref_accessible(): + raise DSLRuntimeError( + "PyIR internal error: re-minting a pyir.ref while the existing " + "cell is still accessible (one-cell-per-place invariant)." + ) + # Mint from the value's OWN raw backing: minting captures this program + # point, never a reload of some attached cell via ``ir_value()``. + ir_val = _raw_backing_ir_value(self._value) + if ir_val is None: + ir_val = self._value.ir_value() + self._ref = pyir.ref(ir_val) + self._ref_context_id = id(ir.Context.current) + + def _is_ref_accessible(self) -> bool: + """True if the existing ref is accessible from the current insertion + point (same or non-isolated ancestor region).""" + if self._ref is None: + return False + if self._ref_context_id != id(ir.Context.current): + return False + current_block = ir.InsertionPoint.current.block + return bool(pyir.value_dominates_ip(self._ref, current_block, None)) + + def _reconstruct(self, loaded_ir: Any) -> Any: + """Reconstruct a DSL value from a loaded MLIR value, preserving compound + wrapper metadata (dtype/shape) via the ``_wrap_ir_like`` template.""" + # ``_wrap_ir_like`` uses ``self._value`` as a template, so compound wrappers + # (e.g. TensorSSA) are rebuilt with full state. + return _wrap_ir_like(self._value, loaded_ir) + + def load(self) -> Any: + """Emit ``pyir.load`` and return a fresh, load-version-tagged DSL value. + ``_mutable_ref`` is deliberately not attached here (callers decide).""" + assert self._ref is not None, ( + "MutableValue.load: no ref -- call take_reference() first" + ) + loaded_ir = pyir.load(self._ref) + self._load_version += 1 + # Record the region-epoch of this choke event (F-GEN): equality with a + # later epoch proves no staged-region boundary was crossed in between. + self._last_choke_region_epoch = _pyir_current_region_epoch() + loaded = self._reconstruct(loaded_ir) + try: + _tag_load( + loaded, + self._load_version, + current_staged_cf_depth(), + _ref_write_epoch(self._ref), + ) + except (AttributeError, TypeError): + pass # value type doesn't accept attrs — no dedup, fine. + return loaded + + def store(self, new_value: Any) -> None: + """Emit ``pyir.store`` of *new_value* into the ref; a pointee-type change + re-mints the cell at the new type and re-books its place row.""" + assert self._ref is not None, ( + "MutableValue.store: no ref -- call take_reference() first" + ) + # ``self._value`` must be set BEFORE ``take_reference`` (which mints from + # it), else the re-mint would use the stale type. + if _pyir_ref_pointee_type_changed(self._ref, new_value): + self._value = new_value + self.take_reference(remint=True) + _pyir_rebook_place_cell(self) + # Store the value's OWN raw backing: a store captures ITS program point, + # never an ``ir_value()`` reload of an attached or foreign cell. + raw = _raw_backing_ir_value(new_value) + if raw is None: + raw = new_value.ir_value() + _place = getattr(self, "_place", None) + _pyir_emit_store( + raw, + self._ref, + var=_place[-1] if isinstance(_place, tuple) and _place else None, + ) + self._value = new_value + # A place dual-booked as a bare-ref row advances its template half too, + # so a row-keyed read reconstructs the class current at this store. + if self._place is not None and self._place in _slot_refs: + _pyir_record_slot_template(self._place, new_value) + # Invalidate previously-loaded values: subsequent uses must reload. + self._load_version += 1 + self._store_version += 1 + # Record the region-epoch of this choke event (F-GEN): equality with a + # later epoch proves no staged-region boundary was crossed in between. + self._last_choke_region_epoch = _pyir_current_region_epoch() + # Stamp write-epoch + staged-CF depth + region-entry clock so a later + # read can judge currency; the load-version tag is NOT stamped + # (presence means "from a pyir.load"). + try: + object.__setattr__( + new_value, _PYIR_REF_EPOCH_ATTR, _ref_write_epoch(self._ref) + ) + object.__setattr__( + new_value, _PYIR_LOAD_DEPTH_ATTR, current_staged_cf_depth() + ) + object.__setattr__( + new_value, _PYIR_REGION_ENTRY_ATTR, _pyir_region_entry_clock() + ) + _pyir_note_holder_write(new_value) # one stamp covers the cluster + except (AttributeError, TypeError): + pass # value type doesn't accept attrs -- no staleness tracking. + + @property + def ref(self) -> ir.Value | None: + """The raw ``pyir.ref`` SSA value (or ``None``).""" + return self._ref + + def __repr__(self) -> str: + return f"MutableValue({self._value!r})" + + +def _ref_dominates_whole_function(mv: "MutableValue") -> bool: + """True if *mv*'s ref is defined in the function entry block (only such a ref + may be published for post-region loads). Conservative: False.""" + if pyir is None: + return False + ref = getattr(mv, "_ref", None) + if ref is None: + return False + try: + entry_block = _get_function_entry_block() + if entry_block is None: + return False + defining_op = ref.owner # OpResult -> defining op (pyir.ref is an op) + op = getattr(defining_op, "operation", defining_op) + return op.block == entry_block + except Exception: + return False + + +def _get_instance_attrs(obj: object) -> list[str]: + """Instance attribute names from instance storage (``__dict__`` plus set + ``__slots__``); skips dunders, never class-level properties/methods.""" + items = _instance_storage_items(obj) + if items is None: + return [] + return [name for name in items if not name.startswith("__")] + + +# Per-class cache of the value-tree protocol judgment. The protocol is +# class-declared (both methods live on the type; instance storage never carries +# them — grep-gated), so a class verdict is stable. Classes with dynamic +# attribute hooks keep the exact per-instance check (bucket "instance"). +_VT_PROTOCOL_CLASS_CACHE: "dict[type, object]" = {} + + +def _vt_protocol_class_bucket(cls: type, obj: object) -> "bool | str": + """Classify *cls* for the protocol check: ``True`` (declares it), + ``False`` (cannot provide it), or ``"instance"`` (hooks; per-instance).""" + declared = any( + "__extract_mlir_values__" in k.__dict__ for k in cls.__mro__ + ) and any("__new_from_mlir_values__" in k.__dict__ for k in cls.__mro__) + if declared: + # Confirm through the instance view once (descriptor shape). + if callable(getattr(obj, "__extract_mlir_values__", None)) and callable( + getattr(obj, "__new_from_mlir_values__", None) + ): + return True + return "instance" + # Undeclared: a dynamic hook could still fabricate the methods. + if ( + type(cls) is type + and cls.__getattribute__ is object.__getattribute__ + and getattr(cls, "__getattr__", None) is None + ): + return False + return "instance" + + +def _implements_dynamic_expression(obj: object) -> bool: + """True if *obj* exposes the DSL ``DynamicExpression`` protocol; its + non-extracted fields are constant context that must not be carried. + + Declaring BOTH dunders is the documented public opt-in (typing.py): only + such classes are ever reconstructed via ``__new_from_mlir_values__``.""" + cls = type(obj) + bucket = _VT_PROTOCOL_CLASS_CACHE.get(cls) + if bucket is None: + bucket = _VT_PROTOCOL_CLASS_CACHE[cls] = _vt_protocol_class_bucket(cls, obj) + if bucket is True: + return True + if bucket is False: + return False + return callable(getattr(obj, "__extract_mlir_values__", None)) and callable( + getattr(obj, "__new_from_mlir_values__", None) + ) + + +def _rebuild_tuple_like(template: tuple, elements: Any) -> tuple: + """Rebuild a tuple-shaped value PRESERVING the template's subtype; an + unrebuildable subclass refuses loudly rather than degrade to plain tuple.""" + elems = tuple(elements) + make = getattr(type(template), "_make", None) + if make is not None and hasattr(template, "_fields"): + try: + return make(elems) + except Exception as exc: + raise DSLUserCodeError( + DiagId.CONTAINER_TUPLE_SUBCLASS_NOT_REBUILDABLE, + type=type(template).__name__, + detail=f" (rebuilding {len(elems)} item(s) raised {exc!r})", + ) from exc + if type(template) is tuple: + return elems + # Extra per-instance state (an instance __dict__ with entries) cannot be + # reproduced from the elements alone -- refuse rather than drop it. + _extra_state = getattr(template, "__dict__", None) + if _extra_state: + raise DSLUserCodeError( + DiagId.CONTAINER_TUPLE_SUBCLASS_NOT_REBUILDABLE, + type=type(template).__name__, + detail=" (it carries extra per-instance attributes)", + ) + try: + return type(template)(elems) + except Exception as exc: + raise DSLUserCodeError( + DiagId.CONTAINER_TUPLE_SUBCLASS_NOT_REBUILDABLE, + type=type(template).__name__, + detail=f" (its constructor rejected the item iterable: {exc!r})", + ) from exc + + +def _is_compound_single_leaf(value: object) -> bool: + """True if *value* is ref-supported yet compound (holds a staged semantic + sub-field): unsupported as a field of a whole-replaced container.""" + if not (_is_staged_value(value) and _can_create_ref(value)): + return False + if not _has_instance_storage(value): + return False + for attr_name in _get_instance_attrs(value): + # Skip the scalar's own MLIR leaf (``value``) and PyIR plumbing attrs; + # only a SEMANTIC staged sub-field makes the value compound. + if ( + attr_name == "value" + or attr_name.startswith("_pyir_") + or (attr_name == "_mutable_ref") + ): + continue + sub = getattr(value, attr_name, None) + if _is_staged_value(sub): + return True + return False + + +def _is_leaf_decomposable(value: object) -> bool: + """True if *value* needs no further decomposition: a meta primitive or a + ref-trackable staged scalar (compound single-leaf values excluded).""" + if value is None or isinstance(value, (int, float, bool, str, bytes)): + return True + if ( + _is_staged_value(value) + and _can_carry_leaf_ref(value) + and not _is_compound_single_leaf(value) + ): + return True + return False + + +def _check_tuple_decomposable( t: tuple, _visited: set[int], ) -> tuple[bool, bool]: - """Check if all elements of a tuple are decomposable. + """Check if all elements of a tuple are decomposable; returns ``(ok, has_staged)``.""" + has_staged = False + for elem in t: + if _is_leaf_decomposable(elem): + if _is_staged_value(elem) and _can_carry_leaf_ref(elem): + has_staged = True + continue + if isinstance(elem, tuple): + ok, found = _check_tuple_decomposable(elem, _visited) + if not ok: + return False, False + has_staged = has_staged or found + continue + # Compound single-leaf staged value (e.g. _Tensor): not decomposable as + # a tuple element -- see _check_all_fields_decomposable for rationale. + if _is_staged_value(elem): + return False, False + if _has_instance_storage(elem) and _check_all_fields_decomposable( + elem, _visited=_visited + ): + has_staged = True + continue + return False, False + return True, has_staged + + +def _check_all_fields_decomposable( + obj: object, + *, + _visited: set[int], +) -> bool: + """True if every instance attribute is a meta primitive, a ref-compatible + staged leaf, or a tuple/nested compound of same; at least one staged.""" + obj_id = id(obj) + if obj_id in _visited: + return False + _visited.add(obj_id) + + if not _has_instance_storage(obj): + return False + + attrs = _get_instance_attrs(obj) + if not attrs: + return False + + has_any_staged = False + for attr_name in attrs: + value = getattr(obj, attr_name) + + if _is_leaf_decomposable(value): + if _is_staged_value(value) and _can_carry_leaf_ref(value): + has_any_staged = True + continue + + if isinstance(value, tuple): + ok, found_staged = _check_tuple_decomposable(value, _visited) + if not ok: + return False + has_any_staged = has_any_staged or found_staged + continue + + if isinstance(value, dict): + # A dict field decomposes per entry when every value is a decomposable + # leaf (the tuple rule's dict analogue); nested containers refuse. + _dict_staged = False + for _dv in list(dict.values(value)): + if _is_leaf_decomposable(_dv) and not isinstance( + _dv, (tuple, list, dict) + ): + if _is_staged_value(_dv) and _can_carry_leaf_ref(_dv): + _dict_staged = True + continue + return False + has_any_staged = has_any_staged or _dict_staged + continue + + # A ref-supported COMPOUND single-leaf field is not decomposable (its + # per-field ref does not survive loop iter_args): refuse loudly instead. + if _is_staged_value(value): + return False + + if _has_instance_storage(value) and _check_all_fields_decomposable( + value, _visited=_visited + ): + has_any_staged = True + continue + + if _implements_dynamic_expression(obj): + # *obj* declares its carried state via the value-tree protocol; a field outside it + # (e.g. a pipeline/barrier ref) is constant context and does not block decomposition. + continue + + return False + + return has_any_staged + + +def _flatten_tuple(t: tuple) -> Any: # Generator[Any, None, None] + """Yield all non-tuple leaf elements from a (possibly nested) tuple.""" + for elem in t: + if isinstance(elem, tuple): + yield from _flatten_tuple(elem) + else: + yield elem + + +# pyir_assign / pyir_read — AST-inserted hooks + + +def _pyir_assign_simple(owner: Any, key: Any, value: Any) -> Any: + """Register a slot in ``_SLOT_REGISTRY`` and return *value* unchanged; + emits no IR, so it is usable without an active MLIR context.""" + _pyir_register_candidate_holder(owner) + _pyir_register_candidate_holder(value) + slot_id = _make_slot_id(owner, key) + mv = _SLOT_REGISTRY.get(slot_id) + if mv is None: + # Placeholder MutableValue (no ``pyir.ref`` yet -- callers may have no MLIR context); it gives the + # registry the right identity and is upgraded when staged CF actually needs a ref. + mv = MutableValue.__new__(MutableValue) + _pyir_setattr_raw(mv, "_value", value) + _pyir_setattr_raw(mv, "_type", type(value)) + _pyir_setattr_raw(mv, "_ref", None) + _pyir_setattr_raw(mv, "_ref_context_id", None) + _pyir_setattr_raw(mv, "_load_version", 0) + _pyir_setattr_raw(mv, "_store_version", 0) + _pyir_setattr_raw(mv, "_last_choke_region_epoch", None) + _SLOT_REGISTRY[slot_id] = mv + return value + + +def _pyir_record_fresh_object_leaf_first_defs(ctor: Any, obj: Any) -> None: + """Record in-region first-def facts for the fields of a directly-constructed + *obj*: the binding position becomes each field's birth block (F-BIRTHPOS).""" + if pyir is None or is_inside_constexpr_loop(): + return + if not (isinstance(ctor, type) and type(ctor) is type): + return + # Freshness: a plain-__new__ class, or exactly the stdlib SimpleNamespace + # (its distinct C-slot __new__ allocates fresh state all the same). + if ( + ctor.__new__ is not object.__new__ and ctor is not types.SimpleNamespace + ) or type(obj) is not ctor: + return + if _is_staged_value(obj): + return + try: + block = ir.InsertionPoint.current.block + except Exception: + return + depth = current_staged_cf_depth() + _record_fresh_leaf_first_defs_walk(obj, block, depth, {id(obj)}) + + +def _subobject_is_ctor_private(parent: Any, value: Any) -> bool: + """Whether a ctor-born *parent*'s field holds a PRIVATE fresh sub-object: + a plain-``__new__`` instance whose only referrer is the parent's own + ``__dict__``. Such a sub-object is reborn with the parent on every + execution of the binding, so its leaves inherit the parent's birth facts + (F-BIRTHPOS). Any other referrer may be an outer alias smuggled into the + ctor (``self.sub = OUTER``) whose state genuinely carries -- do not + descend; the pre-existing conservative carry stays. + + Design tradeoff: this is an after-the-fact ``gc.get_referrers`` heap scan + (trace-time only, behind cheap short-circuit gates) because the binds + that decide privacy happen in plain ``__init__`` bodies and module code, + which never reach the assign choke -- there is no bind-time record to + consult. If ctor-internal binds ever become choke-visible, replace this + with bind-time ownership bookkeeping.""" + if _is_staged_value(value): + return False + cls = type(value) + if not (isinstance(cls, type) and cls.__new__ is object.__new__): + return False + if _implements_dynamic_expression(value): + return False + # A ``__slots__`` parent has no ``__dict__``: its hold then shows up as + # the parent INSTANCE itself among the referrers, which the loop below + # accepts, so a missing dict is not disqualifying. + parent_dict = getattr(parent, "__dict__", None) + # No instance state (dict or set slots) means no leaves at any depth: + # skip the heap scan. Uses the storage view the walker enumerates. + if not _instance_storage_items(value): + return False + # The trace's own keepalive registries hold traced values against GC; + # they are bookkeeping, not user aliases, so they don't disqualify + # (every STRONG-holding keepalive self-registers in + # ``_PYIR_TRACE_KEEPALIVES``; weakref-valued tables such as + # ``_OWNER_KEEPALIVE`` never appear among referrers at all). + # The parent's hold shows up as its ``__dict__`` or as the parent + # INSTANCE itself (managed dicts, ``__slots__``); accept both -- and + # measured on 3.10 and 3.12, the calling frame's locals never appear + # as referrers, so the verdict does not depend on frame lifetime. + saw_parent = False + for ref in gc.get_referrers(value): + if ref is parent_dict or ref is parent: + saw_parent = True + continue + if any(ref is reg for reg in _PYIR_TRACE_KEEPALIVES): + continue + return False + return saw_parent + + +def _record_fresh_leaf_first_defs_walk( + obj: Any, block: "ir.Block", depth: int, seen: "set[int]" +) -> None: + """Record the leaf first-def facts of one ctor-born object and descend + into its ctor-private sub-objects (``t.sub.n``-style depth-2+ leaves).""" + # Meta-primitive fields of value-tree-protocol objects belong to the leaf + # machinery, not to the attr-place ledger. + record_meta_fields = not _implements_dynamic_expression(obj) + for attr_name in _get_instance_attrs(obj): + value = getattr(obj, attr_name, None) + # F-BIRTHPOS depth totality: every published instance-attr place of a + # ctor-born object carries the birth-depth fact (first-wins). + _any_slot = _make_slot_key(None, obj, attr_name) + if _any_slot is not None and _any_slot not in _slot_first_def_depth_any: + _slot_first_def_depth_any[_any_slot] = depth + # Staged scalar leaf: record ONLY the birth block; a later cell mint for + # this place seeds at that block instead of a fabricated entry init. + if _is_staged_value(value) and _can_carry_leaf_ref(value): + leaf_slot = _make_slot_key(None, obj, attr_name) + if leaf_slot is not None and leaf_slot not in _slot_refs: + _slot_first_def_block[leaf_slot] = block + _pyir_keepalive_generation_obj(obj) + continue + if not record_meta_fields: + continue + if type(value) not in (bool, int, float): + if id(value) not in seen and _subobject_is_ctor_private(obj, value): + seen.add(id(value)) + _record_fresh_leaf_first_defs_walk(value, block, depth, seen) + continue + leaf_slot = _make_slot_key(None, obj, attr_name) + if leaf_slot is None or leaf_slot in _slot_refs: + continue + _slot_first_def_inside_cf[leaf_slot] = True + _slot_first_def_depth[leaf_slot] = depth + _slot_first_def_block[leaf_slot] = block + + +def _record_meta_primitive_first_def( + target_name: str, + value: Any, + owner: Any, + slot_name: Any, + maybe_written: bool = True, +) -> Any: + """Record D1 first-def bookkeeping for a Python-primitive *value* (tuples + recurse per leaf); *maybe_written=False* (bound once) skips the wrap.""" + if isinstance(value, tuple): + # Recurse per leaf with the ``f"{name}[{i}]"`` convention shared with ``_decompose_tuple``: + # bare-name tuples use a name-based key; slot-backed tuples carry their parent slot via ``elem_slot``. + has_slot_ctx = owner is not None and slot_name is not None + wrapped = [] + for i, elem in enumerate(value): + elem_name = f"{target_name}[{i}]" + elem_slot = _place_seg_child(slot_name, i) if has_slot_ctx else None + wrapped.append( + _record_meta_primitive_first_def( + elem_name, + elem, + owner if has_slot_ctx else None, + elem_slot, + maybe_written=maybe_written, + ) + ) + # Identity-preserving when no leaf changed (an alias bind keeps + # `t2 is t1` Python-true); else preserve a namedtuple's subtype on + # first-def: ``tuple(wrapped)`` would strip its field names. + if all(w is e for w, e in zip(wrapped, value)): + return value + return _rebuild_tuple_like(value, wrapped) + + # F-CEPLACE: a synthesized tuple-leaf name is its own binding birth (the + # assign choke only saw the whole-tuple target). + if owner is None and slot_name is None and isinstance(target_name, str): + _ce_note_local_assign(target_name) + + # Constexpr-loop per-iteration reset: a raw-primitive first-def into an + # already-promoted slot stores into the existing ref (no staged-CF depth here). + if ( + is_inside_constexpr_loop() + and not is_inside_staged_cf() + and type(value) in (bool, int, float) + and pyir is not None + ): + _ce_slot = _make_slot_key(target_name, owner, slot_name) + if _ce_slot is not None: + _ce_ref = _slot_refs.get(_ce_slot) + if _ce_ref is not None: + log().info( + "[pyir_assign] '%s' constexpr-loop per-iteration reset into " + "already-promoted slot -> store + load", + target_name, + ) + _pyir_emit_store(_emit_constant_for_ref(_ce_ref, value), _ce_ref) + return _load_as_dsl(_ce_ref, place=_ce_slot) + + # Staged scalar leaf first-def (F-BIRTHPOS): record the binding block so a + # later lazy cell mint for this place seeds at the birth position. + if ( + is_inside_staged_cf() + and _is_staged_value(value) + and _can_carry_leaf_ref(value) + and pyir is not None + ): + birth_slot = _make_slot_key(target_name, owner, slot_name) + if birth_slot is not None and birth_slot not in _slot_refs: + try: + _slot_first_def_block[birth_slot] = ir.InsertionPoint.current.block + except Exception: + _slot_first_def_block.pop(birth_slot, None) + + # Wrap a RAW Python literal (exactly bool/int/float, not an already-``_WatchedM`` value) first-read + # inside staged CF, or re-materialise its reset. The strict ``type`` test excludes slot-tracked wrappers. + if is_inside_staged_cf() and type(value) in (bool, int, float): + d1_slot = _make_slot_key(target_name, owner, slot_name) + if d1_slot is not None: + # Per-iteration reset: an earlier unrolled iteration already promoted this slot, so this + # first-def is the body-top reset -- store + ``_load_as_dsl`` so the local carries the reset SSA. + existing_ref = _slot_refs.get(d1_slot) + if existing_ref is not None and pyir is not None: + log().info( + "[pyir_assign] '%s' Python-primitive first-def into " + "already-promoted slot -> store + load (per-iteration reset)", + target_name, + ) + new_ir = _emit_constant_for_ref(existing_ref, value) + _pyir_emit_store(new_ir, existing_ref) + return _load_as_dsl(existing_ref, place=d1_slot) + # Read-only first-def gate (symmetric to ``_pyir_read_impl``): a primitive proved bound once + # (``maybe_written=False``) skips the ``_WatchedM`` wrap, keeping the read a constexpr. + if not maybe_written and not _meta_uses.get(d1_slot): + log().info( + "[pyir_assign] '%s' read-only primitive first-def " + "(maybe_written=False) -> meta passthrough", + target_name, + ) + else: + value = _WatchedM(value, d1_slot) + + # Record whether this slot was first-defined inside staged CF (a later reassignment consults it for the + # per-iteration reset). Gated on ``isinstance`` so ``_WatchedM`` records. + if isinstance(value, (bool, int, float)): + d1_slot = _make_slot_key(target_name, owner, slot_name) + if d1_slot is not None: + inside_cf = is_inside_staged_cf() + _slot_first_def_inside_cf[d1_slot] = inside_cf + if inside_cf: + _slot_first_def_depth[d1_slot] = current_staged_cf_depth() + if pyir is not None: + try: + _slot_first_def_block[d1_slot] = ir.InsertionPoint.current.block + except Exception: + _slot_first_def_block.pop(d1_slot, None) + return value + + +# Base-form machinery: the watched-meta wrappers, dominance/constant helpers and +# slot lookups the entrypoints/carry spine drives. + + +def _cached_ir_value_dominates_current_ip(value: "ir.Value") -> bool: + """Whether *value* is reachable from ``InsertionPoint.current`` (reusing a + sibling-region constant would be a region escape). True when undecidable.""" + if pyir is None: + return True + try: + current_block = ir.InsertionPoint.current.block + except (RuntimeError, ValueError): + return True + return bool(pyir.value_dominates_ip(value, current_block, None)) + + +def _mlir_type_or_none(value: object) -> "ir.Type | None": + """Return *value*'s MLIR type, or ``None`` if it has no SSA backing. + + Reads the baked backing structurally first, never through ``ir_value()`` + (a full staged-read choke that can emit a refresh ``pyir.load`` or bake a + constant -- pure waste when only the type is consumed). A value with no + baked backing keeps the ``ir_value()`` judgment (a meta payload bakes and + reports its type).""" + raw = _raw_backing_ir_value(value) + if raw is not None: + try: + return raw.type + except Exception: + return None + try: + return value.ir_value().type # type: ignore[attr-defined] + except Exception: + return None + + +def _staged_type_changed(old_value: object, new_value: object) -> bool: + """Return True when a same-name reassignment changes the MLIR type. + + A ref's element type is fixed at creation: a same-name rebind to a different + MLIR type would round-trip the OLD type, so the path mints a fresh ref. + + Without both MLIR types the comparison reports "no change" so the + ref-reuse path is preserved. + """ + old_ty = _mlir_type_or_none(old_value) + new_ty = _mlir_type_or_none(new_value) + if old_ty is None or new_ty is None: + return False + return old_ty != new_ty + + +def _pyir_lookup_slot_from_value(value: "Any") -> "MutableValue | None": + """Find the ``MutableValue`` that most recently produced *value* via ``.load()``. + + Cheap path: ``value._mutable_ref`` only; ``None`` when not slot-backed + (callers then route through the plain-Python path with no auto-load). + + Narrow contract (no registry-scan fallback): O(1) lookups and no ambiguity + when two slots share the same ``_load_version``. + """ + if _get_load_version(value) is None: + return None + return getattr(value, "_mutable_ref", None) + + +def _pyir_value_tracked_by_accessible_ref(value: "Any") -> bool: + """Whether *value* is backed by a slot whose ``pyir.ref`` is reachable + from the current insertion point. + + Used by the op-build dominance check to skip values that the lowering + carries through ``scf`` iter_args (the pass maintains SSA dominance). + """ + mv = _pyir_lookup_slot_from_value(value) + if mv is None: + return False + if getattr(mv, "ref", None) is None: + return False + return mv._is_ref_accessible() + + +def _pyir_row_load_type_mismatch(ref: "ir.Value", current_value: Any) -> bool: + """True when *ref*'s pointee differs from *current_value*'s baked MLIR + type: a place row is single-typed, so a differently-typed binding belongs + to another generation of the place and the row must not serve it (the + load would read a stale sibling value). A binding with NO baked backing + (a meta primitive reading its promoted row) only refuses an opaque row.""" + try: + pointee = ref.type.pointee + except Exception: + return False + raw = _raw_backing_ir_value(current_value) + if raw is None: + return _pyir_ir_type_is_opaque(pointee) + try: + return raw.type != pointee + except Exception: + return False + + +def _load_as_dsl( + ref: "ir.Value", + *, + attach: bool = True, + place: Any = None, + stamp_place: bool = False, +) -> Any: + """Emit ``pyir.load(ref)`` and reconstruct the row's wrapper from its + store-time template (F-TYPEID) through the one ``_wrap_ir_like`` funnel. + + *place* keys the template of the row being loaded; *stamp_place* also + stamps the synthetic cell handle with it (F-PLACE) so a later read can + validate the route by place equality. Only the place-NAMED read chokes + stamp: a wrapper returned from a write choke flows into machinery carry + temps whose reads must still claim the cell by first-wins rooting. + + Attaches a synthetic ``_mutable_ref`` so ``_pyir_auto_load_arg`` re-emits + ``pyir.load`` at op boundaries (post-CF uses would carry stale SSA). + + *attach=False* returns a snapshot pinned at this program point: no later + choke re-follows the cell (statement-top read-anchor semantics). + """ + if pyir is None: + return None + loaded_ir = pyir.load(ref) + template = _slot_templates.get(place) if place is not None else None + dsl_val = _wrap_ir_like(template, loaded_ir) + if dsl_val is loaded_ir: + # Identity reconstruction (no wrapper template): nothing to attach to. + return dsl_val + + if not attach: + return dsl_val + # Synthetic MutableValue so _pyir_auto_load_arg can re-load. + try: + mv = MutableValue(dsl_val) + mv._ref = ref + mv._ref_context_id = id(ir.Context.current) + # F-PLACE: the synthetic cell handle names the place it loads from, + # so a later read can validate this route by place equality. + if stamp_place and place is not None: + mv._place = place + # The load just emitted IS this row's choke event (F-GEN). + mv._last_choke_region_epoch = _pyir_current_region_epoch() + _attach_mutable_ref(dsl_val, mv, "D1 _load_as_dsl") + # Freshness tags: this value IS a load of the cell at this program point, + # provably current (same version/depth/epoch + dominance) -- no reload needed. + _tag_load( + dsl_val, + mv._load_version, + current_staged_cf_depth(), + _ref_write_epoch(ref), + ) + except Exception: + pass # attach failure is non-fatal; value still works inside CF + return dsl_val + + +def _clear_slot_mv(owner: Any, slot_name: Any) -> None: + """Remove the recorded ``MutableValue`` for ``(owner, slot_name)``; + silently does nothing when the slot has no entry.""" + if owner is None: + return + store = _slot_store_for_tier1(owner) + if store is not None: + store.pop(slot_name, None) + # Registry-slot owners (no ``__dict__``) live only here; tier-1 owners may + # also have a same-object echo that must not resurface on the next lookup. + _SLOT_REGISTRY.pop(_make_slot_id(owner, slot_name), None) + + +def _pyir_retire_place_row(target_name: Any, owner: Any, slot_name: Any) -> None: + """A fresh-generation rebind abandons the place's previous ledger row; the + binding's next read or write mints the new generation's cell (W2). + + The retire is total across scopes: inside a constexpr instance the + (instance-qualified) key names the same stale row, and leaving it live + would let R1 place authority serve the abandoned generation.""" + if owner is not None and slot_name is not None: + _clear_slot_mv(owner, slot_name) + return + if not isinstance(target_name, str): + return + try: + place = _make_slot_key(target_name, None, None) + except Exception: + return + if place is not None: + _PLACE_REGISTRY.pop(place, None) + + +def _pyir_region_fresh_raw(value: "Any") -> "Any | None": + """Clause-B freshness for a routed raw consumed ACROSS a region boundary: + a fresh load of the cell when the cell is unchanged since the value's own + load (epoch-equal) but a staged region opened after it -- the runtime + re-entry re-executes later-traced stores the pre-region snapshot cannot + see. ``None`` keeps the captured raw: an epoch-MISMATCHED raw is a LAW-1 + retained snapshot (the tuple-swap pin relies on it), and no route or no + region-entry proof keeps the conservative raw.""" + mv = getattr(value, "_mutable_ref", None) + if mv is None or pyir is None or not is_pyir_enabled(): + return None + try: + if mv.ref is None or not mv._is_ref_accessible(): + return None + entry_tag = _get_load_region_entry(value) + if ( + entry_tag is not None + and entry_tag < _pyir_region_entry_watermark() + and _get_load_epoch(value) == _ref_write_epoch(mv.ref) + ): + return mv.load() + except Exception: + return None + return None + + +def _pyir_refresh_cell_read(value: "Any") -> "Any | None": + """The staged-read choke (clause B read half): return a FRESH wrapper + loaded from *value*'s place cell when its cached raw is no longer the + cell's current content, or ``None`` to keep the cached raw. + + ``Numeric.to(ir.Value)`` is the single funnel every staged scalar read + materialises through; a read of the place must observe the cell. + + Staleness discrimination matches ``_pyir_auto_load_arg`` (version + + staged-CF depth + ref write-epoch + dominance). + + Machine stores consume the value's OWN raw and never route through this + choke. Conservative: any failure returns ``None`` (cached-raw behaviour). + """ + mv = getattr(value, "_mutable_ref", None) + if mv is None: + if pyir is None or not is_pyir_enabled(): + return None + return _pyir_recover_place_cell_read(value) + if pyir is None or not is_pyir_enabled(): + return None + try: + # A cell minted in a DIFFERENT MLIR context belongs to a previous + # compilation; any dereference is a use-after-free. Refuse loudly. + if mv._ref is not None and mv._ref_context_id != id(ir.Context.current): + raise DSLUserCodeError( + "This value was produced by a PREVIOUS @jit compilation " + "(its backing IR lives in a different, already-finalized " + "compilation context) and cannot be reused here.", + suggestion=( + "Pass the value through the kernel's arguments (a " + "runtime value) or recompute it in this compilation; " + "values stashed across separate top-level @jit calls " + "do not carry over." + ), + ) + if mv.ref is None or not mv._is_ref_accessible(): + return None + # Pairing-identity split: full freshness discrimination only when *value* + # is the cell's CANONICAL representative (last stored/loaded value). + + # A site-less read of an aliased wrapper keeps its dominating raw (it + # cannot tell which place it names); a trapped raw loads the attached cell. + if mv._value is not value: + if _value_dominates_current_ip(value): + # A literal re-bake would re-materialize a possibly-unstored + # placeholder seed: a placeholder-seeded cell serves the read + # through a live load the used-poison scan can judge. + if _is_literal_backed(value) and not _pyir_ref_init_is_real_value( + mv.ref + ): + return mv.load() + fresh = _pyir_region_fresh_raw(value) + if fresh is not None: + return fresh + return None + return mv.load() + cached_version = _get_load_version(value) + if ( + cached_version is not None + and cached_version == mv._load_version + and _get_load_depth(value) == current_staged_cf_depth() + and _get_load_epoch(value) == _ref_write_epoch(mv.ref) + and _value_dominates_current_ip(value) + ): + return None + if cached_version is None: + # Version-untagged wrapper (a store stamps epoch+depth, not version): + # its raw is current only at the same epoch AND depth, and dominating. + + # A paired wrapper with NO tags carries no freshness proof: the read + # loads the cell (a redundant load-after-store folds downstream). + epoch_tag = _get_load_epoch(value) + depth_tag = _get_load_depth(value) + if ( + epoch_tag is not None + and epoch_tag == _ref_write_epoch(mv.ref) + and depth_tag is not None + and depth_tag == current_staged_cf_depth() + and _value_dominates_current_ip(value) + ): + return None + return mv.load() + except Exception: + return None + + +def _pyir_holder_pairs_binding(value: "Any", candidates: "Any") -> "list[tuple]": + """The ``(holder, attr)`` pairs among *candidates* whose instance storage + binds *value* by identity.""" + pairs: "list[tuple]" = [] + for cand in candidates: + if not _has_instance_storage(cand): + continue + for attr in _get_instance_attrs(cand): + try: + if getattr(cand, attr, None) is value: + pairs.append((cand, attr)) + except Exception: + continue + return pairs + + +def _pyir_recovery_holder_pairs(value: "Any") -> "Any": + """The holder pairs an unpaired read may resolve (a list, or a lazy + registry-order iterator). Inside an extraction walk the DECLARED owner + context is authoritative: the innermost declared holder binding the value + serves the read, a two-attr binding within it refuses loudly (the walk + names the holder, never the attr), and no context hit means no emission + (the existing loud refusal stands). Outside a walk, the registry scan in + sighting order is the recorded C4 residual (multi-holder states reach + here through channels the walks do not declare yet).""" + if _EXTRACTION_WALK_OWNERS: + for cand in reversed(_EXTRACTION_WALK_OWNERS): + pairs = _pyir_holder_pairs_binding(value, (cand,)) + if len(pairs) == 1: + return pairs + if len(pairs) > 1: + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.SNAPSHOT_UNMATERIALIZABLE, + filename=filename, + lineno=lineno, + var=type(value).__name__, + ) + return [] + return _pyir_registry_pairs_iter(value) + + +def _pyir_registry_pairs_iter(value: "Any") -> "Any": + """Registry-order pair enumeration with the full scan's membership and + order. Non-funneled candidates are probed EAGERLY up front (their + ``getattr`` can run user descriptor getters, which the full scan always + executed); funneled wrappers are probed lazily -- those probes are plain + storage reads (frame-filtered recorder aside), so probes past the + consumer's first accepted pair are droppable with no observable + difference.""" + cands = _pyir_registry_candidate_objects() + eager: "dict[int, list]" = {} + for i, cand in enumerate(cands): + if not _pyir_wrapper_write_funneled(type(cand)): + eager[i] = _pyir_holder_pairs_binding(value, (cand,)) + + def _pairs() -> "Any": + for i, cand in enumerate(cands): + pre = eager.get(i) + if pre is not None: + yield from pre + continue + if not _has_instance_storage(cand): + continue + for attr in _get_instance_attrs(cand): + try: + if getattr(cand, attr, None) is value: + yield (cand, attr) + except Exception: + continue + + return _pairs() + + +def _pyir_recover_place_cell_read(value: "Any") -> "Any | None": + """Owner-context recovery for an unpaired staged read: promote / reload + from the resolved holder's live cell; refuse a type-differing in-region + rebind loudly.""" + if _is_literal_backed(value) and is_inside_staged_cf(): + if getattr(value, "_pyir_place_probe_neg", False): + return None + if not (_is_staged_value(value) and _can_create_ref(value)): + return None + for cand, attr in _pyir_recovery_holder_pairs(value): + slot_mv = _get_slot_mv(cand, attr) + if slot_mv is not None and slot_mv._is_ref_accessible(): + if _pyir_ref_pointee_type_changed(slot_mv._ref, value): + return None + loaded = slot_mv.load() + else: + # One cell per place: a bare-ref row already booked for this + # place (a promoted meta slot) is the cell -- adopt it. + place = _make_slot_key(None, cand, attr) + d1_ref = _slot_refs.get(place) if place is not None else None + if d1_ref is not None: + if _pyir_ref_pointee_type_changed(d1_ref, value): + return None + loaded = _load_as_dsl(d1_ref, place=place) + try: + setattr(cand, attr, loaded) + except Exception: + pass + return loaded + slot_mv = _create_ref(value) + slot_mv = _set_slot_mv(cand, attr, slot_mv) + loaded = slot_mv.load() + try: + setattr(cand, attr, loaded) + except Exception: + pass + return loaded + try: + _pyir_setattr_raw(value, "_pyir_place_probe_neg", True) + except (AttributeError, TypeError): + pass + return None + try: + raw = _raw_backing_ir_value(value) + if raw is None or _value_dominates_current_ip(value): + return None + except Exception: + return None + assert raw is not None # narrowed above; the except path returned + for cand, attr in _pyir_recovery_holder_pairs(value): + slot_mv = _get_slot_mv(cand, attr) + if slot_mv is None or not slot_mv._is_ref_accessible(): + continue + slot_ref = slot_mv._ref + if slot_ref is not None and _pyir_ref_pointee_type_changed(slot_ref, value): + _pyir_raise_type_changed_in_region(attr, slot_ref.type.pointee, raw.type) + loaded = slot_mv.load() + try: + setattr(cand, attr, loaded) + except Exception: + pass + return loaded + return None + + +def _pyir_meta_primitive_value(v: "Any") -> "Any": + """Return the Python value of a meta-valued binding (primitive, ``_WatchedM``, + or cell-less literal-backed staged scalar), else ``_NO_CONST_VALUE``.""" + if isinstance(v, _WatchedM): + return v.python_value + if isinstance(v, (bool, int, float)): + return v + try: + if ( + _is_staged_value(v) + and _is_literal_backed(v) + and getattr(v, "_mutable_ref", None) is None + ): + return v.value + except Exception: + pass + return _NO_CONST_VALUE + + +def _pyir_snapshot_registry_meta_slots() -> "list[tuple[Any, str, Any]]": + """Snapshot ``(holder, attr, literal)`` for registered holder attributes bound + to cell-less literal-backed staged scalars at a staged-loop entry.""" + out: "list[tuple[Any, str, Any]]" = [] + if not is_pyir_enabled(): + return out + for cand in _pyir_registry_candidate_objects(): + if isinstance(cand, MutableValue) or _is_staged_value(cand): + continue + if not _has_instance_storage(cand): + continue + for attr in _get_instance_attrs(cand): + if isinstance(attr, str) and attr.startswith("_pyir"): + continue + try: + v = getattr(cand, attr, None) + except Exception: + continue + try: + if ( + v is not None + and _is_staged_value(v) + and _is_literal_backed(v) + and getattr(v, "_mutable_ref", None) is None + ): + out.append((cand, attr, v.value)) + except Exception: + continue + return out + + +def _pyir_verify_registry_meta_slots( + snapshot: "list[tuple[Any, str, Any]]", +) -> None: + """Refuse a meta-valued holder attribute advanced inside a staged loop body + outside the ledger (traced once -- only the first iteration's value bakes).""" + for holder, attr, pre in snapshot: + try: + cur = getattr(holder, attr, None) + except Exception: + continue + curp = _pyir_meta_primitive_value(cur) + if curp is _NO_CONST_VALUE or isinstance(cur, _WatchedM): + # The cell's literal init is the witness: minted in-region from a + # value differing from the loop-entry literal = an unledgered advance. + + # A non-literal init (an SSA seed) is a genuine carried store -- exempt. + mv = getattr(cur, "_mutable_ref", None) + if mv is None or mv.ref is None: + continue + try: + init_const = _ref_literal_init_const(mv.ref) + except Exception: + continue + if init_const is _NO_CONST_VALUE: + continue + try: + if bool(init_const == pre): + continue + except Exception: + continue + else: + try: + unchanged = bool(curp == pre) + except Exception: + unchanged = True + if unchanged: + continue + d1_slot = _make_slot_key(None, holder, attr) + if d1_slot is not None and d1_slot in _slot_refs: + continue + raise DSLUserCodeError( + DiagId.BOUNDARY_META_LOOP_CARRY, + name=str(attr), + ) + + +# Opaque-owner subscript-write audit: a subscript store through a type outside +# the place-owner families (dict/list, watched, memref-like, staged / +# DynamicExpression) runs a ``__setitem__`` the chokes cannot see into. +# Pre-existing reachable leaves are carried by the region walks; a container +# key CREATED by such a store has no pre-region cell, so no lowering can make +# it follow the region's runtime predicate. The write choke arms this audit +# with a key-set snapshot of the owner's reachable containers; the region +# close verifies and retires it. +_PYIR_OPAQUE_OWNER_AUDITS: "dict[int, tuple[str, Any, list[tuple[Any, frozenset]]]]" = {} + + +def _pyir_opaque_owner_key_sets(owner: "Any") -> "list[tuple[Any, frozenset]]": + """Key sets of every dict/list reachable from *owner* through instance + storage and container values (attribute-walk only: no protocol method of + the owner runs, so the snapshot itself has no side effects).""" + out: "list[tuple[Any, frozenset]]" = [] + seen: "set[int]" = set() + + def _walk(obj: "Any") -> None: + if obj is None or id(obj) in seen: + return + seen.add(id(obj)) + if isinstance(obj, dict): + out.append((obj, frozenset(dict.keys(obj)))) + for v in list(dict.values(obj)): + _walk(v) + return + if isinstance(obj, list): + out.append((obj, frozenset(range(list.__len__(obj))))) + for v in list.__iter__(obj): + _walk(v) + return + if isinstance(obj, tuple): + for v in obj: + _walk(v) + return + if _is_staged_value(obj): + return + items = _instance_storage_items(obj) + if items is None: + return + for v in list(items.values()): + _walk(v) + + _walk(owner) + return out + + +def _pyir_arm_opaque_owner_audit(name: str, owner: "Any") -> None: + """Arm the region-close key-set audit for an opaque-``__setitem__`` owner. + The first write's pre-state wins: the audit diffs the state before the + owner's first in-region store against the region close.""" + if id(owner) in _PYIR_OPAQUE_OWNER_AUDITS: + return + _PYIR_OPAQUE_OWNER_AUDITS[id(owner)] = ( + name, + owner, + _pyir_opaque_owner_key_sets(owner), + ) + + +def _pyir_verify_opaque_owner_audits() -> None: + """Refuse (loudly) any container key created through an opaque + ``__setitem__`` owner inside dynamic staged CF; retires the armed audits + whether or not they pass (each dynamic region judges its own writes).""" + if not _PYIR_OPAQUE_OWNER_AUDITS: + return + try: + for name, owner, snap in list(_PYIR_OPAQUE_OWNER_AUDITS.values()): + pre_by_id = {id(h): ks for h, ks in snap} + for holder, cur_keys in _pyir_opaque_owner_key_sets(owner): + pre_keys = pre_by_id.get(id(holder), frozenset()) + added = [k for k in cur_keys if k not in pre_keys] + if added: + raise DSLUserCodeError( + DiagId.CONTAINER_OPAQUE_SUBSCRIPT_KEY_CREATED, + var=name, + kind=type(owner).__name__, + detail=", ".join(sorted(repr(k) for k in added)), + ) + finally: + _PYIR_OPAQUE_OWNER_AUDITS.clear() + + +def _pyir_leaf_is_carried(value: "Any") -> bool: + """Return True if a Numeric leaf is carried through PyIR slot state. + + An instrumented-method mutation leaves a slot marker (slot ref / + ``_mutable_ref`` / ``_pyir_load_version``); an uninstrumented + (``@dsl_user_op``) mutation rebinds the parent to an untracked Numeric. + + A carried leaf keeps its markers even while its SSA is region-trapped. + Conservative: any error returns True (an indeterminate leaf is never clobbered). + """ + try: + if _pyir_value_tracked_by_accessible_ref(value): + return True + if getattr(value, "_mutable_ref", None) is not None: + return True + if _get_load_version(value) is not None: + return True + return False + except Exception: + return True + + +def _ref_lives_in_strictly_enclosing_block(value: "Any") -> bool: + """True if *value*'s backing ``pyir.ref`` was minted in a block that + STRICTLY ENCLOSES the current insertion point. + + A conditional update of an outer-scope binding stores back through the + shallower ref; a flat same-region redefinition takes a fresh shadow. + + Block nesting is the reliable signal (the depth dict and a live + ``_mutable_ref`` cannot discriminate). Any error returns ``False``. + """ + if pyir is None: + return False + mv = getattr(value, "_mutable_ref", None) + if mv is None: + return False + ref_ssa = getattr(mv, "ref", None) + if ref_ssa is None: + return False + try: + owner = ref_ssa.owner + ref_block = ( + owner + if isinstance(owner, ir.Block) + else getattr(owner, "operation", owner).block + ) + cur_block = ir.InsertionPoint.current.block + if ref_block == cur_block: + return False + return pyir.is_value_in_ancestor_region(ref_ssa, cur_block) + except Exception: + return False + + +def _carries_foreign_slot_binding(value: Any) -> bool: + """Return True if *value* (or any of its tuple leaves) re-resolves a + storage slot on every use rather than holding a fixed SSA value. + + Re-resolution = an accessible ``_mutable_ref`` (auto-load re-emits + ``pyir.load`` per consumer) or a ``_slot_refs``-promoted ``_WatchedM`` key. + + A live binding is correct for the slot's own accessor but WRONG for an + independent Python local; see ``_freeze_foreign_slot_binding``. + + Containers recurse per element: a literal / comprehension / concat + container SNAPSHOTS its elements in Python value terms. + """ + if isinstance(value, tuple): + return any(_carries_foreign_slot_binding(e) for e in value) + if isinstance(value, list): + return any(_carries_foreign_slot_binding(e) for e in value) + if isinstance(value, dict): + return any(_carries_foreign_slot_binding(e) for e in value.values()) + if isinstance(value, _WatchedM): + slot = value._slot_key + return slot is not None and slot in _slot_refs + mv = getattr(value, "_mutable_ref", None) + return mv is not None and mv._is_ref_accessible() + + +def _freeze_foreign_slot_binding( + value: Any, context: str, *, fresh_container: bool = False +) -> Any: + """Return a copy of *value* frozen to its current SSA, detached from + any storage slot it was loaded from. + + Used by ``pyir_assign`` for a fresh plain-local binding that still + re-resolves a foreign slot: freezing binds it to its assignment-time value. + + Value-preserving: reuses the already-loaded SSA and only removes the slot + linkage. Tuples are frozen leaf-by-leaf; unfreezable values pass through. + + A ``list`` / ``dict`` is frozen element-in-place ONLY under the declared + ``fresh_binding`` AST fact (constructed by the binding statement itself). + + An aliased container is returned untouched -- detaching a shared list's + elements would rewrite the source binding. Top-level snapshot only. + """ + if isinstance(value, tuple): + # Shape-faithful freeze (clause A3): rebuild through the tuple's own + # subtype so a namedtuple binding keeps its named fields. + return _rebuild_tuple_like( + value, (_freeze_foreign_slot_binding(e, context) for e in value) + ) + if isinstance(value, list): + if fresh_container: + for i, e in enumerate(value): + if isinstance(e, (list, dict)): + continue # possibly-aliased nested container + if _carries_foreign_slot_binding(e): + value[i] = _freeze_foreign_slot_binding(e, context) + return value + if isinstance(value, dict): + if fresh_container: + for k, e in list(value.items()): + if isinstance(e, (list, dict)): + continue # possibly-aliased nested container + if _carries_foreign_slot_binding(e): + value[k] = _freeze_foreign_slot_binding(e, context) + return value + if isinstance(value, _WatchedM): + slot = value._slot_key + if slot is None or slot not in _slot_refs: + return value + # Bake/reuse the current SSA, then return a slot-less wrapper so + # ``_pyir_auto_load_arg`` treats it as a fixed value. + frozen = _WatchedM(value.python_value, slot_key=None) + try: + frozen._cached_ir = value.ir_value() + except Exception: + return value + log().info("[pyir_assign] froze _WatchedM binding (%s)", context) + return frozen + mv = getattr(value, "_mutable_ref", None) + if mv is None or not mv._is_ref_accessible(): + return value + # A SUPERSEDED handle of an admitted rebound memref cell must not freeze + # to the slot's current value (that is a fresh generation's buffer). + _pyir_refuse_stale_memref_serve(value, mv) + # Capture the slot's CURRENT carried value at assignment time, then detach. A leaf still on its + # construction-time SSA would bake a STALE pre-loop value, so load the ref (iter_arg) and freeze THAT. + try: + frozen = mv.load() + except Exception: + frozen = _fresh_wrapper(value) + if frozen is value: + return value + # Defensive: neither ``load`` nor the ``_fresh_wrapper`` fallback copies + # ``_mutable_ref``; ensure it is detached. + if getattr(frozen, "_mutable_ref", None) is not None: + try: + _pyir_delattr_raw(frozen, "_mutable_ref") + except (AttributeError, TypeError): + pass + log().info("[pyir_assign] froze ref-backed binding (%s)", context) + return frozen + + +def _pyir_same_ref_reload(value: Any, ref: "ir.Value | None") -> bool: + """Shared same-ref reload exemption: *value* is backed by *ref*'s own cell. + + A value whose ``_mutable_ref`` IS the slot's ref is a reload of (or an + instrumented store through) the SAME place cell, not an unledgered advance. + """ + if ref is None: + return False + try: + return getattr(getattr(value, "_mutable_ref", None), "ref", None) is ref + except Exception: + return False + + +def _is_opaque_leaf_type(value: "ir.Value") -> bool: + """Return ``True`` for an OPAQUE value-tree leaf -- one the scalar + ``pyir.ref`` / M2S slot machinery does NOT auto-carry. + + Opaque = a dialect-specific value (an MMA / copy atom handle etc.) the + slot machinery cannot carry; the region carry carries it instead. + + Gates the ``scf.if`` forward-carry to opaque leaves only; a slot-carried + scalar leaf is never double-carried. Any error -> ``False``. + """ + try: + return _pyir_ir_type_is_opaque(value.type) + except Exception: + return False + + +def _pyir_ir_type_is_opaque(ty: Any) -> bool: + """``True`` for a dialect-opaque MLIR type (not a builtin integer / index / + float / shaped type). Conservative: any error -> ``False``.""" + try: + if ir.IntegerType.isinstance(ty): + return False + if ir.IndexType.isinstance(ty): + return False + if hasattr(ir, "FloatType") and ir.FloatType.isinstance(ty): + return False + if ir.ShapedType.isinstance(ty): + # vector<...> / tensor<...> of scalars -- carried by the slot + # machinery as a staged aggregate. + return False + return True + except Exception: + return False + + +def _pyir_extract_leaf_values(obj: Any) -> "list[ir.Value] | None": + """Return *obj*'s canonical value-tree leaves as bare ``ir.Value``s. + + Uses ``__extract_mlir_values__`` (the non-PyIR pytree flattening), so the + result is the independent leaf set that reconstructs *obj*. + + Narrower than :func:`_pyir_walk_ir_value_holders`: a whole-object rebind + carries only canonical leaves; the reconstruct re-derives derived fields. + + Returns ``None`` for a non-value-tree object or any extraction failure. + + A SCALAR leaf is normalised to an ``ArithValue`` via + :func:`_as_arith_capable_scalar_leaf`; an OPAQUE leaf is left raw. + """ + if not _implements_dynamic_expression(obj): + return None + try: + vals = obj.__extract_mlir_values__() + except Exception: + return None + if not isinstance(vals, (list, tuple)): + return None + leaves: "list[ir.Value]" = [] + for v in vals: + if isinstance(v, ir.Value): + leaves.append(_as_arith_capable_scalar_leaf(v)) + else: + raw = _raw_backing_ir_value(v) + if raw is None: + return None + leaves.append(_as_arith_capable_scalar_leaf(raw)) + return leaves + + +def _pyir_is_carryable_tuple(obj: Any) -> bool: + """A TOP-LEVEL ``tuple`` / ``list`` whose every element is a value-tree: the + region-close carry treats it as a value-tree (concat of leaves).""" + return ( + isinstance(obj, (tuple, list)) + and len(obj) > 0 + and all(_implements_dynamic_expression(e) for e in obj) + ) + + +def _pyir_extract_any_leaf_values(obj: Any) -> "list[ir.Value] | None": + """:func:`_pyir_extract_leaf_values` extended to a top-level tuple / list: + the canonical leaves are the concatenation of each element's, in order. + + Falls back to the plain value-tree extract for a non-container; ``None`` + on any non-extractable element.""" + if _pyir_is_carryable_tuple(obj): + out: "list[ir.Value]" = [] + for elem in obj: + ev = _pyir_extract_leaf_values(elem) + if ev is None: + return None + out.extend(ev) + return out + return _pyir_extract_leaf_values(obj) + + +def _pyir_new_from_mlir_values_any(obj: Any, new_vals: "list[ir.Value]") -> Any: + """``__new_from_mlir_values__`` extended to a top-level tuple / list: split + *new_vals* by extract arity and rebuild a NEW container of the same kind. + + Falls back to the plain protocol for a non-container; ``None`` on any + failure (the caller leaves the binding unchanged). + + Opens the rebuild-protocol frame (F-BIRTH): token adoption is legal only + for objects this machinery produces (V-6), and the carried source must + still be its born class (V-7).""" + _PYIR_REBUILD_PROTOCOL_DEPTH[0] += 1 + try: + if _pyir_is_carryable_tuple(obj): + rebuilt: "list[Any]" = [] + idx = 0 + for elem in obj: + _pyir_validate_owner_class(elem) + try: + arity = len(elem.__extract_mlir_values__()) + except Exception: + return None + new_elem = elem.__new_from_mlir_values__(new_vals[idx : idx + arity]) + # Re-root the rebuilt element's place at the source's owner token. + _pyir_propagate_owner_token(elem, new_elem) + rebuilt.append(new_elem) + idx += arity + # Shape-faithful container rebuild (clause A3): a namedtuple / tuple + # subclass must come back as its own subtype (fields as separate args). + if isinstance(obj, list): + new_container: Any = list(rebuilt) + else: + new_container = _rebuild_tuple_like(obj, rebuilt) + _pyir_propagate_owner_token(obj, new_container) + return new_container + _pyir_validate_owner_class(obj) + new_obj = obj.__new_from_mlir_values__(new_vals) + # Re-root the rebuilt object's place at the source's owner token. + _pyir_propagate_owner_token(obj, new_obj) + return new_obj + finally: + _PYIR_REBUILD_PROTOCOL_DEPTH[0] -= 1 + + +def _is_opaque_leaf_value_tree(obj: Any) -> bool: + """Return ``True`` for a value-tree object whose canonical extract-leaves + are ALL opaque (a tensor's opaque memref handle, an MMA / copy atom handle). + + Its loop-carried state is carried by the opaque-leaf region carry, + not the scalar slot machinery (which cannot carry an opaque leaf). + + A whole-object rebind of such an object PASSES THROUGH in ``pyir_assign`` + and leaves the carry to the region-close carry. + + Keys on the CANONICAL leaves (what ``__new_from_mlir_values__`` + reconstructs from), unlike a raw ``__dict__`` walk over derived fields. + + Conservative: non-value-tree, no leaves, or any scalar leaf -> ``False`` + (a scalar value-tree stays on the M->M decompose path). + """ + leaves = _pyir_extract_leaf_values(obj) + if not leaves: + return False + return all(_is_opaque_leaf_type(leaf) for leaf in leaves) + + +def _value_trees_different_prototype(old: object, new: object) -> bool: + """Return True if *old* and *new* are value-tree-protocol objects of the same + class whose canonical leaves ALIGN (same count) but a NON-LEAF *prototype* + field (a meta-primitive or a primitive tuple) DIFFERS. + + That is the signature of a whole-object BRANCH-SELECT between two + structurally distinct prototypes of one class. + + Such a rebind carries only the DIFFERING canonical leaves through the + region-close ``scf.if``; the reconstruct restores the prototype fields. + + A ``__dict__`` walk would instead branch-select every field, tripping the + meta-primitive guard and leaving a ref unstored on the no-op branch. + + A SAME-prototype carry returns False (stays on the ``__dict__``-walk + decomposition). Conservative: any failure / non-protocol -> False. + """ + if old is new: + return False + if not _implements_dynamic_expression(old) or not _implements_dynamic_expression( + new + ): + return False + if type(old) is not type(new): + return False + if not _has_instance_storage(old) or not _has_instance_storage(new): + return False + try: + old_leaves = _pyir_extract_leaf_values(old) or [] + new_leaves = _pyir_extract_leaf_values(new) or [] + except Exception: + return False + # Leaves must align for region-close carry + reconstruct; a divergent + # extract is left to the conservative ``__dict__`` walk. + if len(old_leaves) == 0 or len(old_leaves) != len(new_leaves): + return False + try: + attrs = _get_instance_attrs(old) + except Exception: + return False + for attr_name in attrs: + try: + ov = getattr(old, attr_name) + nv = getattr(new, attr_name) + except AttributeError: + continue + # A non-leaf prototype field is a meta-primitive or tuple of primitives; + # an SSA-backed (staged) field is a leaf handled by carry. + if ov is None or isinstance(ov, (bool, int, float)): + if (nv is None or isinstance(nv, (bool, int, float))) and ov != nv: + return True + elif ( + isinstance(ov, tuple) + and isinstance(nv, tuple) + and all(isinstance(x, (bool, int, float)) for x in ov) + and all(isinstance(x, (bool, int, float)) for x in nv) + and ov != nv + ): + return True + return False + + +def _has_accessible_loop_carried_ref(target_name: str, value: object) -> bool: + """Return True iff *value* already has a ``pyir.ref`` slot reachable from + inside the current loop body -- so its top-of-body read would load the + loop-carried value rather than bind to the pre-loop SSA. + + Used by :func:`pyir_promote_loop_body_arg` to decide whether a body-entry + ``pyir_read`` is needed; checks ``_mutable_ref`` and the D1 ``_slot_refs``. + + A ref counts only when accessible from the current IP. Conservative on + error: report "no ref" so the body-entry load is emitted. + """ + try: + mv = getattr(value, "_mutable_ref", None) + if mv is not None and mv._is_ref_accessible(): + return True + d1_slot = _make_slot_key(target_name, None, None) + if d1_slot is not None and d1_slot in _slot_refs: + return True + except Exception: + return False + return False + + +def _pyir_skip_trapped_child_region_writeback( + ref: "ir.Value", + new_value: Any, + place: Any = None, + mv: "MutableValue | None" = None, +) -> Any: + """Shared guard for the two ``pyir_assign`` writeback fast-paths.""" + if pyir is None or ref is None: + return None + if _raw_backing_ir_value(new_value) is None: + return None + if _value_dominates_current_ip(new_value): + return None + if mv is not None: + loaded = mv.load() + _attach_mutable_ref(loaded, mv, "child-region writeback skip") + return loaded + return _load_as_dsl(ref, place=place) + - Returns ``(ok, has_staged)``. +def _walk_value_tree_holders( + obj: Any, + leaf_pred: "Callable[[Any], bool]", + _visited: "set[int] | None" = None, + _out: "list[tuple[Any, str, Any]] | None" = None, +) -> "list[tuple[Any, str, Any]]": + """Return the ``(holder, attr, value)`` of every *leaf_pred* leaf in *obj*. + + Appends into *_out* when given (one shared accumulator across a multi-root + sweep; identical DFS order to the concatenating form) and returns it. """ - has_staged = False - for elem in t: - if _is_leaf_decomposable(elem): - if _is_staged_value(elem) and _can_create_ref(elem): - has_staged = True + if _visited is None: + _visited = set() + if _out is None: + _out = [] + oid = id(obj) + if oid in _visited: + return _out + if not _implements_dynamic_expression(obj): + return _out + storage = _instance_storage_items(obj) + if storage is None: + return _out + _visited.add(oid) + + for attr_name, attr_val in list(storage.items()): + if isinstance(attr_val, (tuple, list)): + for elem in attr_val: + _walk_value_tree_holders(elem, leaf_pred, _visited, _out) + elif isinstance(attr_val, dict): + # Dict fields recurse per value (raw ``dict.values`` so an adopted + # watched dict is walked without firing its read chokes). + for elem in list(dict.values(attr_val)): + _walk_value_tree_holders(elem, leaf_pred, _visited, _out) + elif leaf_pred(attr_val): + _out.append((obj, attr_name, attr_val)) + else: + _walk_value_tree_holders(attr_val, leaf_pred, _visited, _out) + return _out + + +# Exact builtin scalar types that can never hold or be a staged-SSA leaf; +# subclasses deliberately excluded (exact-type check) so they take the full path. +_WALK_PRIM_SKIP = frozenset({int, float, bool, str, bytes, complex, type(None)}) + + +def _pyir_walk_ir_value_holders(obj: Any) -> "list[tuple[Any, str, Any]]": + """Return the in-place holders of every staged-SSA leaf in *obj*. + + Specialized form of ``_walk_value_tree_holders`` with the staged-SSA leaf + predicate inlined (identical DFS output; this sweep runs once per staged + region over the whole candidate registry, so it is hot).""" + out: "list[tuple[Any, str, Any]]" = [] + _walk_ir_value_holders_into(obj, set(), out) + return out + + +# Branch class of one attr value in the holder walk, memoized per type +# (class shape is definition-time-stable). Mirrors the walk's branch +# precedence exactly: containers, then bare ``ir.Value``, then exact builtin +# scalars, then the generic leaf-or-recurse path. +_WALK_KIND_SEQ, _WALK_KIND_DICT, _WALK_KIND_IRV, _WALK_KIND_PRIM, _WALK_KIND_OTHER = ( + 1, + 2, + 3, + 4, + 5, +) +_WALK_CLASS_KIND: "dict[type, int]" = {} + + +def _walk_class_kind(cls: type) -> int: + if issubclass(cls, (tuple, list)): + return _WALK_KIND_SEQ + if issubclass(cls, dict): + return _WALK_KIND_DICT + if issubclass(cls, ir.Value): + return _WALK_KIND_IRV + if cls in _WALK_PRIM_SKIP: + return _WALK_KIND_PRIM + return _WALK_KIND_OTHER + + +def _walk_storage_prim_irv_only(storage: "dict[str, Any]") -> bool: + """True when every storage value is an exact builtin scalar or a bare + ``ir.Value`` -- on such nodes the walk's loop body is append/continue only + (no user code reachable), so it may iterate the live dict directly.""" + kinds = _WALK_CLASS_KIND + for v in storage.values(): + cls = type(v) + kind = kinds.get(cls) + if kind is None: + kind = kinds[cls] = _walk_class_kind(cls) + if kind != _WALK_KIND_PRIM and kind != _WALK_KIND_IRV: + return False + return True + + +def _walk_ir_value_holders_into( + obj: Any, + visited: "set[int]", + out: "list[tuple[Any, str, Any]]", +) -> None: + """Recursive body of :func:`_pyir_walk_ir_value_holders`.""" + oid = id(obj) + if oid in visited: + return + if not _implements_dynamic_expression(obj): + return + storage = _instance_storage_items(obj) + if storage is None: + return + visited.add(oid) + kinds = _WALK_CLASS_KIND + if _walk_storage_prim_irv_only(storage): + # PRIM/IRV-only node: the loop body runs no user code, so no snapshot + # copy of the storage dict is needed (same records, same order). + for attr_name, attr_val in storage.items(): + if kinds[type(attr_val)] == _WALK_KIND_IRV: + out.append((obj, attr_name, attr_val)) + return + for attr_name, attr_val in list(storage.items()): + cls = type(attr_val) + kind = kinds.get(cls) + if kind is None: + kind = kinds[cls] = _walk_class_kind(cls) + if kind == _WALK_KIND_OTHER: # scalar-Numeric leaf or recurse + if _is_numeric_leaf_holder(attr_val): + out.append((obj, attr_name, attr_val)) + else: + _walk_ir_value_holders_into(attr_val, visited, out) + elif kind == _WALK_KIND_IRV: # bare ir.Value leaf + out.append((obj, attr_name, attr_val)) + elif kind == _WALK_KIND_PRIM: # exact builtin scalar: nothing below continue - if isinstance(elem, tuple): - ok, found = _check_tuple_decomposable(elem, _visited) - if not ok: - return False, False - has_staged = has_staged or found + elif kind == _WALK_KIND_SEQ: + for elem in attr_val: + _walk_ir_value_holders_into(elem, visited, out) + else: + # Raw ``dict.values`` so an adopted watched dict is walked without + # firing its read chokes. + for elem in list(dict.values(attr_val)): + _walk_ir_value_holders_into(elem, visited, out) + + +def _pyir_build_gather_segment(cand: Any) -> "tuple | None": + """The walk records of a cacheable sweep root as ``(records, guards)``, or + ``None`` when the root must take the live walk. ``records`` holds the + ``(attr, value)`` pairs; ``guards`` holds ``(wrapper, stamp)`` pairs for + every Numeric-leaf record, revalidated at splice time. The root itself is + left out of the cached row (the splice re-attaches it): a cached record + holding its own root would pin the root -- and with it the root's registry + row -- for the whole trace, defeating the registry's weakref pruning. + + Cacheable = write-funneled class (declared value-tree protocol, no + ``__slots__``/override) whose every storage value is an exact builtin + scalar, a bare ``ir.Value``, an SSA-backed scalar-Numeric leaf, or a + walk-inert value (protocol bucket ``False``: the walk's recursion into it + returns before visiting). Rebinding any such field goes through the + declared write funnel and stamps the root, so the records are a pure + function of root storage. The one non-root input is the SSA-backed leaf + classification, which can only change through a write on the WRAPPER's own + funnel; the per-record guard pins the wrapper's write stamp so a stale + classification invalidates the row instead of splicing. + This builder mirrors the live walk on such roots, record for record.""" + cls = type(cand) + if _VT_PROTOCOL_CLASS_CACHE.get(cls) is not True: + return None # unseen/instance/False buckets: the live walk judges it + if not _pyir_wrapper_write_funneled(cls): + return None + storage = _instance_storage_items(cand) + if storage is None: + return None + kinds = _WALK_CLASS_KIND + stamps = _PYIR_HOLDER_WRITE_STAMPS + seg: "list[tuple[str, Any]]" = [] + guards: "list[tuple[Any, int]]" = [] + for attr_name, attr_val in storage.items(): + vcls = type(attr_val) + kind = kinds.get(vcls) + if kind is None: + kind = kinds[vcls] = _walk_class_kind(vcls) + if kind == _WALK_KIND_IRV: # bare ir.Value leaf record + seg.append((attr_name, attr_val)) + elif kind == _WALK_KIND_PRIM: # exact builtin scalar: nothing below continue - # Compound single-leaf staged value (e.g. _Tensor): not decomposable as - # a tuple element -- see _check_all_fields_decomposable for rationale. - if _is_staged_value(elem): - return False, False - if hasattr(elem, "__dict__") and _check_all_fields_decomposable( - elem, _visited=_visited - ): - has_staged = True + elif kind == _WALK_KIND_OTHER and _is_numeric_leaf_holder(attr_val): + # SSA-backed scalar-Numeric leaf: the live walk records the + # wrapper by reference without recursing. An unregistered + # wrapper has no stamp row to guard on -- live walk. + wstamp = stamps.get(id(attr_val)) + if wstamp is None: + return None + seg.append((attr_name, attr_val)) + guards.append((attr_val, wstamp)) + elif kind == _WALK_KIND_OTHER and _VT_PROTOCOL_CLASS_CACHE.get(vcls) is False: + # Class-stable walk-inert value (e.g. the slot cell handle): the + # walk's recursion returns before visiting, so no record. continue - return False, False - return True, has_staged + else: # containers, walkable objects, meta-payload wrappers: live walk + return None + return (tuple(seg), tuple(guards)) -def _check_all_fields_decomposable( - obj: object, - *, - _visited: set[int], +class _PyirSelfLeafHolder: + """Sentinel ``holder`` of a SELF-LEAF snapshot record -- the record's TYPE + discriminator (recognised by :func:`_is_self_leaf_record`).""" + + +_PYIR_SELF_LEAF_HOLDER = _PyirSelfLeafHolder() + + +def _is_self_leaf_record(holder: Any, attr_name: str) -> bool: + """Return ``True`` for a SELF-LEAF snapshot record (structural: the holder is the + dedicated :class:`_PyirSelfLeafHolder` sentinel).""" + return isinstance(holder, _PyirSelfLeafHolder) + + +def _self_leaf_snapshot(arg: Any) -> "list[tuple[Any, str, ir.Value]] | None": + """Snapshot a WHOLE-OBJECT bare ``ir.Value`` *arg* as ONE self-leaf record.""" + if not isinstance(arg, ir.Value): + return None + raw = _raw_backing_ir_value(arg) + if raw is None: + return None + if _implements_dynamic_expression(arg): + # A protocol-bearing ir.Value whose holder walk found nothing is + # still its own single leaf when extraction agrees (e.g. an + # ir.Value-subclass pointer wrapper with no instance storage); + # bail only when extraction names a different leaf set. + try: + leaves = _pyir_extract_any_leaf_values(arg) or [] + except Exception: + return None + if len(leaves) != 1 or not _same_ir_value(leaves[0], raw): + return None + return [(_PYIR_SELF_LEAF_HOLDER, _PYIR_SELF_LEAF_ATTR, raw)] + + +def _numeric_leaf_kind(value: object) -> "_Literal['none', 'ssa', 'meta']": + """Classify *value* as a value-tree scalar-Numeric leaf and its backing.""" + global _IS_STAGED_VALUE_FN + try: + fn = _IS_STAGED_VALUE_FN + if fn is None: + from .multi_stage_manager import _is_staged_value + + fn = _IS_STAGED_VALUE_FN = _is_staged_value + + if isinstance(value, ir.Value): + return _NUMERIC_LEAF_NONE + if not (fn(value) and _can_carry_leaf_ref(value)): + return _NUMERIC_LEAF_NONE + if _raw_backing_ir_value(value) is not None: + return _NUMERIC_LEAF_SSA + return _NUMERIC_LEAF_META + except Exception: + return _NUMERIC_LEAF_NONE + + +def _is_numeric_leaf_holder(value: object) -> bool: + """Return True if *value* is a staged scalar Numeric wrapper backed by a + real SSA (the parent-attribute leaf shape the holder walk snapshots).""" + return _numeric_leaf_kind(value) == _NUMERIC_LEAF_SSA + + +def _pyir_record_promoted_place_leaf(owner: Any, key: Any) -> None: + """Record an owner place promoted at the D1 write choke so the stale-leaf + region reload can repair its raw storage; trace-scoped, weakref preferred.""" + if owner is None or not isinstance(key, (str, int)): + return + oid = id(owner) + row = _PYIR_PROMOTED_PLACE_LEAVES.get(oid) + if row is None: + + def _drop_leaf_row(_r: Any, _oid: int = oid) -> None: + _PYIR_PROMOTED_PLACE_LEAVES.pop(_oid, None) + + try: + handle: Any = _weakref.ref(owner, _drop_leaf_row) + except TypeError: + handle = owner # non-weakrefable (dict/list): strong, cleared at exit + row = (handle, set()) + _PYIR_PROMOTED_PLACE_LEAVES[oid] = row + row[1].add(key) + + +def _pyir_reload_one_stale_leaf(owner: Any, key: Any, context: str) -> None: + """Re-bind ONE owner place to a fresh dominating load when its current + staged scalar escaped a closed region while its slot ref dominates.""" + try: + leaf_place = _make_slot_key(None, owner, key) + ref = _slot_refs.get(leaf_place) + except Exception: + return + if not isinstance(ref, ir.Value) or not _value_dominates_current_ip(ref): + return + is_container = isinstance(owner, (dict, list)) + try: + if is_container: + value = ( + dict.__getitem__(owner, key) + if isinstance(owner, dict) + else list.__getitem__(owner, key) + ) + else: + value = getattr(owner, key) + except Exception: + return + cur_raw = _raw_backing_ir_value(value) + # Only an ESCAPING leaf: a real backing SSA that does NOT dominate the IP. + if cur_raw is None or _value_dominates_current_ip(cur_raw): + return + try: + fresh = _load_as_dsl(ref, place=leaf_place) + if isinstance(owner, dict): + dict.__setitem__(owner, key, fresh) + elif isinstance(owner, list): + list.__setitem__(owner, key, fresh) + else: + _pyir_setattr_raw(owner, key, fresh) + log().info( + "[pyir region] re-loaded stale staged leaf '%s' from its slot ref (%s)", + key, + context, + ) + except Exception: + return + + +def _pyir_reload_stale_staged_attr_leaves(context: str) -> None: + """Re-load any TRACKED object-attribute slot whose current staged scalar + escaped a sibling region (its backing SSA no longer dominates the IP).""" + if pyir is None: + return + for sid in list(_SLOT_REGISTRY.keys()): + if sid.kind != "attr" or not isinstance(sid.key, str): + continue + wref = _OWNER_KEEPALIVE.get(sid.owner) + owner = wref() if wref is not None else None + if owner is None: + continue + # The D1 table keys by PLACE, so the owner object (not the id-keyed + # ``_SlotId``) is needed to resolve the entry. + _pyir_reload_one_stale_leaf(owner, sid.key, context) + # Places promoted at the D1 write choke (no registry row): same repair. + for oid, (handle, keys) in list(_PYIR_PROMOTED_PLACE_LEAVES.items()): + owner = handle() if isinstance(handle, _weakref.ReferenceType) else handle + if owner is None: + continue + for key in list(keys): + _pyir_reload_one_stale_leaf(owner, key, context) + + +def _pyir_revert_escaped_leaf( + holder: Any, + attr_name: str, + snap_value: Any, + allow_numeric_revert: bool = True, ) -> bool: - """Return True if every instance attribute is a meta primitive, - a staged+ref-compatible leaf, a tuple of same, or a nested compound - of same. At least one field must be staged. - """ - obj_id = id(obj) - if obj_id in _visited: + """Revert one snapshotted leaf if it escaped a closed staged region.""" + try: + cur_value = getattr(holder, attr_name) + except AttributeError: return False - _visited.add(obj_id) - - if not hasattr(obj, "__dict__"): + if cur_value is snap_value: return False - - attrs = _get_instance_attrs(obj) - if not attrs: + cur_raw = _raw_backing_ir_value(cur_value) + snap_raw = _raw_backing_ir_value(snap_value) + if cur_raw is None or snap_raw is None: + return False + # A carried Numeric leaf (instrumented mutator -> slot markers) is a genuine carried accumulator; + # reverting it would drop the carry. Only the untracked ``@dsl_user_op`` escape is repaired here. + if _is_numeric_leaf_holder(snap_value) and ( + not allow_numeric_revert or _pyir_leaf_is_carried(cur_value) + ): + return False + # Keep the replacement if it still dominates the outer IP (a legitimate forward update); revert only + # when the replacement is trapped in the closed region AND the snapshot SSA dominates. + if _value_dominates_current_ip(cur_raw): + return False + if not _value_dominates_current_ip(snap_raw): + return False + try: + _pyir_setattr_raw(holder, attr_name, snap_value) + except (AttributeError, TypeError): return False + return True - has_any_staged = False - for attr_name in attrs: - value = getattr(obj, attr_name) - if _is_leaf_decomposable(value): - if _is_staged_value(value) and _can_create_ref(value): - has_any_staged = True +def _pyir_repair_captured_escaped_leaves( + snapshot: "list[tuple[Any, str, Any]] | None", + context: str, + allow_numeric_revert: bool = True, +) -> None: + """Revert captured value-tree leaves that escaped a closed staged region. Companion + to :func:`_pyir_gather_captured_leaf_holders`. + + Records of write-funneled holders unwritten since the snapshot's gather + clock are skipped: unchanged stamp on a funneled class whose ``getattr`` + is a storage read implies exactly the unchanged-identity ``continue``.""" + if snapshot is None: + return + snap_clock = getattr(snapshot, "pyir_gather_clock", None) + stamps = _PYIR_HOLDER_WRITE_STAMPS + for holder, attr_name, snap_value in snapshot: + if snap_clock is not None: + stamp = stamps.get(id(holder)) + if ( + stamp is not None + and stamp <= snap_clock + and _pyir_getattr_reads_storage(type(holder), attr_name) + ): + continue + # Unchanged-identity fast path: mirrors the first check inside + # :func:`_pyir_revert_escaped_leaf`, hoisted out of the call. + try: + if getattr(holder, attr_name) is snap_value: + continue + except AttributeError: continue + if _pyir_revert_escaped_leaf( + holder, attr_name, snap_value, allow_numeric_revert + ): + log().info( + "[pyir region] reverted escaped captured value-tree leaf '%s' (%s)", + attr_name, + context, + ) - if isinstance(value, tuple): - ok, found_staged = _check_tuple_decomposable(value, _visited) - if not ok: - return False - has_any_staged = has_any_staged or found_staged + +def _pyir_repair_region_escaped_leaves( + arg: Any, + snapshot: "list[tuple[Any, str, Any]] | None", + context: str, + allow_numeric_revert: bool = True, +) -> Any: + """Revert value-tree leaves that escaped a closed staged region, in place.""" + if snapshot is None: + return arg + for holder, attr_name, snap_value in snapshot: + # Unchanged-identity fast path: mirrors the first check inside + # :func:`_pyir_revert_escaped_leaf`, hoisted out of the call. + try: + if getattr(holder, attr_name) is snap_value: + continue + except AttributeError: continue + if _pyir_revert_escaped_leaf( + holder, attr_name, snap_value, allow_numeric_revert + ): + log().info( + "[pyir region] reverted escaped value-tree leaf '%s' (%s)", + attr_name, + context, + ) + return arg - # A ref-supported COMPOUND single-leaf (e.g. ``cute._Tensor``, whose - # ``_iterator`` is itself a staged Pointer) is NOT decomposable as a - # container field: the per-field ref it creates does not survive loop - # iter_args, silently dropping the field's loop-carried value. Bail out - # so the container hits the clean ``CONTAINER_OBJECT_REPLACED`` rejection - # (P89: only a TOP-LEVEL whole-replace of a ``_Tensor`` threads, via the - # S->S path). - if _is_staged_value(value): + +def _pyir_ref_init_is_real_value(ref: "ir.Value") -> bool: + """Return ``True`` if a ``pyir.ref``'s INIT operand is a genuine value, not a + deferred-UB stand-in.""" + try: + if pyir is None or not hasattr(pyir, "ref_init_is_real_value"): return False + return bool(pyir.ref_init_is_real_value(ref)) + except Exception: + return False - if hasattr(value, "__dict__") and _check_all_fields_decomposable( - value, _visited=_visited - ): - has_any_staged = True + +def _pyir_repair_region_escaped_tuple_leaves( + arg: Any, + if_op: "ir.Operation", + context: str, +) -> None: + """Re-bind a TUPLE/LIST-contained scalar value-tree leaf trapped in a closed + STANDALONE ``scf.if`` to a dominating ``pyir.load`` of its slot ref.""" + if pyir is None or arg is None: + return + if not _implements_dynamic_expression(arg): + return + for attr_name in _get_instance_attrs(arg): + attr_val = getattr(arg, attr_name, None) + if not isinstance(attr_val, (tuple, list)): continue + new_elems: "list[Any]" = [] + rebuilt = False + for elem in attr_val: + repaired = _pyir_reload_trapped_tuple_leaf(elem, if_op, context) + if repaired is not None: + new_elems.append(repaired) + rebuilt = True + else: + new_elems.append(elem) + if rebuilt: + # ``type(attr_val)(new_elems)`` breaks a namedtuple (positional + # ``__new__``), so rebuild through the tuple-like helper. + rebuilt_val = ( + new_elems + if isinstance(attr_val, list) + else _rebuild_tuple_like(attr_val, new_elems) + ) + try: + _pyir_setattr_raw(arg, attr_name, rebuilt_val) + except (AttributeError, TypeError): + pass + + +def _pyir_reload_trapped_tuple_leaf( + elem: Any, + if_op: "ir.Operation", + context: str, +) -> Any: + """Repair ONE tuple element trapped in a closed ``scf.if``; see + :func:`_pyir_repair_region_escaped_tuple_leaves`.""" + if not _is_numeric_leaf_holder(elem): + return None + raw = _raw_backing_ir_value(elem) + if raw is None: + return None + # Already dominates the post-``if`` IP -> correctly carried, leave it. + if _value_dominates_current_ip(elem): + return None + # Only repair a leaf genuinely TRAPPED inside the closed ``scf.if``. + if not _ir_value_defined_inside_op(raw, if_op): + return None + mv = getattr(elem, "_mutable_ref", None) + if mv is None: + return None + ref = getattr(mv, "ref", None) + if not isinstance(ref, ir.Value): + return None + if not mv._is_ref_accessible() or not _value_dominates_current_ip(ref): + return None + # Never re-bind to a load of a deferred-UB (poison / placeholder) ref: that would reintroduce a + # SCOPE_READ_NEVER_SET read. A real ``arith`` init passes; a conditional-store placeholder is blocked. + if not _pyir_ref_init_is_real_value(ref): + return None + try: + # Row-authoritative reload: the leaf's own cell reconstructs from its + # store-time template and re-attaches for downstream re-loads. + loaded = mv.load() + _attach_mutable_ref(loaded, mv, "trapped tuple-leaf reload") + except Exception as exc: + log().info("[pyir if] tuple-leaf reload emit failed: %s", exc) + return None + loaded_raw = _raw_backing_ir_value(loaded) + if loaded_raw is not None: + # Mutate the SHARED wrapper in place so a sibling region holding this exact wrapper observes the + # dominating load, then return the freshly re-bound leaf for the rebuilt tuple. + try: + _pyir_setattr_raw(elem, "value", loaded_raw) + except (AttributeError, TypeError): + pass + log().info( + "[pyir if] reloaded trapped tuple-contained scalar leaf to a dominating " + "post-if load (%s)", + context, + ) + return loaded + + +def _pyir_mark_loop_carried_ref(ref: "ir.Value") -> None: + """Stamp the carry-ref marker attribute on *ref*'s defining ``pyir.ref``. See + :data:`_PYIR_LOOP_ITER_ARGS_ATTR`.""" + try: + ref_op = _get_defining_operation(ref) + op = getattr(ref_op, "operation", ref_op) + op.attributes[_PYIR_LOOP_ITER_ARGS_ATTR] = ir.UnitAttr.get() + except Exception: + pass + + +def _pyir_ref_carries_value(ref: "ir.Value", raw: "ir.Value") -> bool: + """True iff *raw* is a value the cell *ref* demonstrably carries: its init + operand, a stored operand of one of its ``pyir.store``s, or a load of it.""" + if pyir is None or ref is None or raw is None: + return False + try: + ref_op = _get_defining_operation(ref) + if _same_ir_value(ref_op.operands[0], raw): + return True + # A load OF this ref: the loaded value is the cell's reaching value + # at the load point by definition. + raw_owner = getattr(raw, "owner", None) + if raw_owner is not None and not isinstance(raw_owner, ir.Block): + raw_op = getattr(raw_owner, "operation", raw_owner) + if str(getattr(raw_op, "name", "")) == "pyir.load" and _same_ir_value( + raw_op.operands[0], ref + ): + return True + # A value some ``pyir.store`` wrote into the cell. + for use in ref.uses: + user = getattr(use, "owner", None) + if user is None: + continue + user_op = getattr(user, "operation", user) + try: + if str(user_op.name) != "pyir.store": + continue + if _same_ir_value(user_op.operands[0], raw): + return True + except Exception: + continue + except Exception: + return False + return False + +def _pyir_ref_dominates_op(ref: "ir.Value", op: "ir.Operation") -> bool: + """True iff *ref* dominates the position immediately before *op* -- the + anchor a region-carry ref needs so the C++ pass can lift it.""" + if pyir is None or ref is None or op is None: + return False + try: + op_norm = getattr(op, "operation", op) + return bool(pyir.value_dominates_ip(ref, op_norm.block, op_norm)) + except Exception: return False - return has_any_staged +def _pyir_local_place_for_name(name: "str | None") -> Any: + """LOCAL place key ``('local', scope_id, name)`` for an executor-known bare + name; ``None`` when no scope is open or inside a constexpr unroll.""" + if not isinstance(name, str) or not name: + return None + try: + if is_inside_constexpr_loop(): + return None + if not _PYIR_SCOPE_STACK: + return None + return _make_slot_key(name, None, None) + except Exception: + return None -def _flatten_tuple(t: tuple) -> Any: # Generator[Any, None, None] - """Yield all non-tuple leaf elements from a (possibly nested) tuple.""" - for elem in t: - if isinstance(elem, tuple): - yield from _flatten_tuple(elem) - else: - yield elem +def _pyir_region_carry_place_for( + holder: Any, attr_name: Any, arg_name: "str | None" +) -> Any: + """Ledger place for a region-forwarder leaf record (clause A totality).""" + if _is_self_leaf_record(holder, attr_name): + return _pyir_local_place_for_name(arg_name) + return None + + +def _pyir_adopt_region_carry_ref( + place: Any, snap_value: "ir.Value", region_op: "ir.Operation" +) -> "tuple[ir.Value | None, bool]": + """Get-or-mint THE ``pyir.ref`` cell for a region-carried leaf place.""" + if place is not None and isinstance(place, tuple) and place and place[0] == "local": + cand = _slot_refs.get(place) + if cand is not None: + try: + type_ok = cand.type.pointee == snap_value.type + except Exception: + type_ok = False + if ( + type_ok + and _pyir_ref_dominates_op(cand, region_op) + and _pyir_ref_carries_value(cand, snap_value) + ): + return cand, False + else: + place = None + try: + # Anchor placement (clause A totality for COMPOUND leaves): hoist the + # mint OUT of the loop body so its init anchors before the region. + anchor = getattr(region_op, "operation", region_op) + cur = anchor + while True: + parent = getattr(cur, "parent", None) + if parent is None: + break + parent = getattr(parent, "operation", parent) + name = str(getattr(parent, "name", "")) + if name == "builtin.module" or getattr(parent, "parent", None) is None: + break + try: + if not pyir.value_dominates_ip(snap_value, parent.block, parent): + break + except Exception: + break + if name in ("scf.while", "scf.for"): + anchor = parent + cur = parent + with ir.InsertionPoint(anchor): + ref = pyir.ref(snap_value) + except Exception: + return None, False + _pyir_mark_loop_carried_ref(ref) + if place is not None: + try: + entry_block = _get_function_entry_block() + region_op_norm = getattr(region_op, "operation", region_op) + if entry_block is not None and region_op_norm.block == entry_block: + _slot_refs[place] = ref + # F-TYPEID: one type identity per place -- adopt the live + # registry row's wrapper template when it describes this + # pointee; a raw-SSA leaf row reconstructs as identity. + template: Any = snap_value + row = _PLACE_REGISTRY.get(place) + if row is not None and row._value is not None: + row_raw = _raw_backing_ir_value(row._value) + if row_raw is not None and row_raw.type == snap_value.type: + template = row._value + _pyir_record_slot_template(place, template) + except Exception: + pass + return ref, True + + +def _pyir_adopt_d1_carry_cell( + target_name: str, raw: "ir.Value" +) -> "MutableValue | None": + """Adopt the name's live D1 row as its loop-carry cell (one-cell-per-place): + body reads and the body's D1 write-throughs must share ONE ref. Adoption + requires the row to demonstrably carry *raw*; ``None`` lets the caller mint.""" + try: + place = _make_slot_key(target_name, None, None) + ref = _slot_refs.get(place) if place is not None else None + type_ok = ref is not None and ref.type.pointee == raw.type + except Exception: + # No provable row (key resolution failed, or the row is not a typed + # pointer cell): minting is the correct conservative outcome. + return None + if not type_ok: + return None + if not _pyir_ref_carries_value(ref, raw): + return None + # Guards passed: exceptions past here are internal errors and must + # surface; a silent mint would recreate the split-cell miscompile. + mv = MutableValue(raw) + mv._ref = ref + mv._ref_context_id = id(ir.Context.current) + if not mv._is_ref_accessible(): + return None + return _set_slot_mv(None, target_name, mv) + + +def _pyir_innermost_enclosing_region_op_at_ip() -> "ir.Operation | None": + """The innermost region-carrying op enclosing the current insertion point, + or ``None`` at function scope or on any resolution failure.""" + try: + block = ir.InsertionPoint.current.block + entry_block = _get_function_entry_block() + if entry_block is not None and block == entry_block: + return None + op = block.owner + if op is None: + return None + return getattr(op, "operation", op) + except Exception: + return None + + +def _pyir_mint_ref_at_value_def(raw: "ir.Value") -> "ir.Value | None": + """Mint ``pyir.ref %raw`` immediately at *raw*'s definition point (after its + defining op; at block begin for a block argument) -- the furthest-out anchor.""" + if pyir is None: + return None + try: + entry_block = _get_function_entry_block() + if entry_block is None: + return None + owner = raw.owner + if isinstance(owner, ir.Block): + if owner != entry_block: + return None + with ir.InsertionPoint.at_block_begin(owner): + return pyir.ref(raw) + def_op = getattr(owner, "operation", owner) + if def_op.block != entry_block: + return None + with ir.InsertionPoint.after(def_op): + return pyir.ref(raw) + except Exception: + return None -# ========================================================================== -# pyir_assign / pyir_read — AST-inserted hooks -# ========================================================================== +def _pyir_single_opaque_extract_leaf(value: Any) -> "ir.Value | None": + """The single canonical extract leaf of *value*, for a wrapper whose raw + backing is not directly discoverable (the leaf nests in an interior impl + object, e.g. Array -> _ArrayImpl -> _base). None unless extraction + yields exactly one ``ir.Value``.""" + try: + leaves = _pyir_extract_any_leaf_values(value) or [] + except Exception: + return None + if len(leaves) == 1 and isinstance(leaves[0], ir.Value): + return leaves[0] + return None -def _pyir_assign_simple(owner: Any, key: Any, value: Any) -> Any: - """Simplified (owner, key, value) entry point for the unified slot - registry. - Registers a slot in ``_slot_mvs`` keyed on - ``_make_slot_key(None, owner, key)`` and returns *value* unchanged. +def _pyir_pair_place_cell_handle(place: Any, ref: "ir.Value", value: Any) -> None: + """Pair *value* with the place cell it was just stored into, so a use of the + wrapper AFTER the region reloads instead of consuming the trapped raw. - This entry point is intentionally side-effect-free with respect to - MLIR IR (no ``pyir.ref`` is emitted), so it can be used outside - ``@cute.jit`` traces -- e.g. by the unit test that runs without an - active MLIR context. When an MLIR context is active, production - callers drive ``pyir_assign`` through the full 5-argument signature. - """ - slot_key = _make_slot_key(None, owner, key) - mv = _slot_mvs.get(slot_key) - if mv is None: - # Register a placeholder MutableValue. We don't (and can't, in - # general) emit a ``pyir.ref`` here because callers may invoke - # this without an active MLIR context. The placeholder gives - # the registry the right identity; production code paths upgrade - # it to a fully-initialised ``MutableValue`` when staged CF - # actually needs a ref. - mv = MutableValue.__new__(MutableValue) - object.__setattr__(mv, "_value", value) - object.__setattr__(mv, "_type", type(value)) - object.__setattr__(mv, "_ref", None) - object.__setattr__(mv, "_ref_context_id", None) - object.__setattr__(mv, "_load_version", 0) - _slot_mvs[slot_key] = mv - return value + A bare local needs no pairing: its post-region reads route through the + instrumented name and reach the cell there. An attr/subscript read hands + the wrapper straight to a ``@dsl_user_op``, so only the attached cell lets + ``_pyir_auto_load_arg`` re-emit the load. A wrapper that already names a + live cell keeps it (one cell per place).""" + if isinstance(place, tuple) and place and place[0] == "local": + return + try: + if getattr(value, "_mutable_ref", None) is not None: + return + mv = MutableValue(value) + mv._ref = ref + mv._ref_context_id = id(ir.Context.current) + mv._place = place + _attach_mutable_ref(value, mv, "opaque place-cell rebind") + except Exception: + return -# --------------------------------------------------------------------------- -# Shared side-effect-free IR-navigation substrate (no callers in this layer). -# Dominance folds into one C++ CAPI crossing (pyir.value_dominates_ip). -# --------------------------------------------------------------------------- +def _pyir_place_opaque_rebind_choke( + target_name: str, place: Any, old_value: Any, new_value: Any +) -> "ir.Value | None": + """Clause-A write choke: a PLACE rebound to a same-type OPAQUE value inside a + staged region writes through that place's ONE cell. Attr/subscript places + cell exactly like locals: the place key already roots on the owner token, so + every alias of the owner names the same cell.""" + if pyir is None or place is None: + return None + kind = place[0] if isinstance(place, tuple) and place else None + if kind != "local" and kind not in ("attr", "subscript"): + return None + try: + if is_inside_constexpr_loop(): + return None + except Exception: + return None + old_raw = _raw_backing_ir_value(old_value) + if not isinstance(old_raw, ir.Value): + old_raw = _pyir_single_opaque_extract_leaf(old_value) + new_raw = _raw_backing_ir_value(new_value) + if not isinstance(new_raw, ir.Value): + new_raw = _pyir_single_opaque_extract_leaf(new_value) + if not isinstance(old_raw, ir.Value) or not isinstance(new_raw, ir.Value): + return None + if _same_ir_value(old_raw, new_raw): + return None + try: + if old_raw.type != new_raw.type: + return None + if not _is_opaque_leaf_type(old_raw): + return None + except Exception: + return None + # An owner-rooted place has no per-iteration carry cell of its own, so a + # NEW value that already dominates the region is some other binding's + # entry value (`a.ptr = b.ptr`): celling it would pin that trace-time + # handle for every iteration while the source advances. A bare local is + # safe -- its loop carry re-serves the source each iteration -- so only + # the owner-rooted kinds decline, leaving the pre-existing loud refusal + # in place. Loud beats silent drift. + if kind != "local": + region = _pyir_innermost_enclosing_region_op_at_ip() + if region is None or not _ir_value_defined_inside_op(new_raw, region): + return None + # CONTINUE-STORE leg. + cand = _slot_refs.get(place) + if cand is not None: + try: + type_ok = cand.type.pointee == new_raw.type + except Exception: + type_ok = False + if ( + type_ok + and _value_dominates_current_ip(cand) + and _pyir_ref_carries_value(cand, old_raw) + ): + _pyir_emit_store( + new_raw, cand, choke=f"opaque place rebind '{target_name}'" + ) + _pyir_record_slot_template(place, new_value) + _pyir_pair_place_cell_handle(place, cand, new_value) + return cand + # A live row that does not carry the old value belongs to an earlier + # generation of the name; the mint gate re-books the row. + region_op = _pyir_innermost_enclosing_region_op_at_ip() + if region_op is None: + return None + if not _value_dominates_current_ip(old_raw): + return None + if _ir_value_defined_inside_op(old_raw, region_op): + return None + # The new value may be trapped INSIDE the region (the classic carry) or + # defined OUTSIDE it (a swap/select between pre-region handles): both + # are loop-carried binding state and cell identically -- without a cell + # an outside->outside swap silently pins the trace-time binding. + # Handles cell-carry uniformly (all spaces, all types); refuse only a + # rebind backed by an IN-REGION allocation (entry-hoisted: aliases one + # buffer). Register-space memref handles defer to the store choke + # below, whose liveness gate admits the iteration-private scratch idiom + # or refuses there; every other alloc-rooted value (pointer, view, + # non-register memref) refuses here -- no declared space fact can prove + # its aliasing faithful. + if _pyir_memref_alloc_rooted_inside_region(new_raw, region_op): + if not _pyir_type_is_register_memref(new_raw.type): + _pyir_raise_memref_inregion_alloc_rebind(target_name) + ref = _pyir_mint_ref_at_value_def(old_raw) + if ref is None: + return None + _pyir_mark_loop_carried_ref(ref) + _slot_refs[place] = ref + _pyir_record_slot_template(place, new_value) + _pyir_emit_store(new_raw, ref, choke=f"opaque place rebind '{target_name}'") + _pyir_pair_place_cell_handle(place, ref, new_value) + log().info( + "[pyir_assign] '%s' opaque rebind → place cell store", + target_name, + ) + return ref -def _is_func_boundary_op(op_name: str) -> bool: - """Return True if *op_name* is a function-like op that owns an SSA body - region where a ``pyir.ref`` may be hosted. +def _pyir_value_rooted_outside_loop(value: "ir.Value", loop_op: "ir.Operation") -> bool: + """Return ``True`` if *value* is (transitively) computed ONLY from values defined + OUTSIDE *loop_op* -- i.e.""" + if pyir is None or not isinstance(value, ir.Value): + return False + try: + return bool( + pyir.value_rooted_outside_loop( + value, getattr(loop_op, "operation", loop_op) + ) + ) + except Exception: + return False - Matches every FunctionOpInterface op by the ``.func`` naming - convention plus the explicit entry ops in - :data:`_NON_DOT_FUNC_ENTRY_OPS`. Name-pattern (not enumeration) so new - dialect function ops need no maintenance here. - """ - return op_name.endswith(".func") or op_name in _NON_DOT_FUNC_ENTRY_OPS +def _pyir_memref_alloc_rooted_inside_region( + value: "ir.Value", region_op: "ir.Operation" +) -> bool: + """True iff *value*'s operand cone, cut at the region boundary, reaches an + op DECLARING the Allocate memory effect INSIDE *region_op* -- the value + depends on memory allocated inside the staged region. Conservative: False.""" + if pyir is None or not isinstance(value, ir.Value): + return False + if not hasattr(pyir, "value_alloc_rooted_inside_region"): + return False + try: + return bool( + pyir.value_alloc_rooted_inside_region( + value, getattr(region_op, "operation", region_op) + ) + ) + except Exception: + return False -def _is_module_boundary_op(op_name: str) -> bool: - """Return True if *op_name* is a module-like symbol-table container. - The entry-block walk stops here: a module owns no SSA region for a - ``pyir.ref``, and walking past it risks dereferencing recycled - top-of-module block wrappers. - """ - return op_name in _MODULE_OPS or op_name.endswith(".module") +def _pyir_value_is_inner_loop_carried( + value: "ir.Value", loop_op: "ir.Operation" +) -> bool: + """Return ``True`` if *value* is a ``pyir.load`` of a leaf-carry ref that a + NESTED loop / if (closed inside *loop_op*) already created.""" + try: + if not isinstance(value, ir.Value): + return False + def_op = _get_defining_operation(value) + if getattr(def_op, "name", None) != "pyir.load": + return False + ref_val = def_op.operands[0] + ref_op = _get_defining_operation(ref_val) + op = getattr(ref_op, "operation", ref_op) + # Stamped-marker consult: the marker is the exact record the carry + # machinery established when it minted the ref. + if _PYIR_LOOP_ITER_ARGS_ATTR not in op.attributes: + return False + # The ref must live INSIDE the current loop body (a nested loop closed within it). A ref at the + # same scope as the current loop (its own outer ref from a prior pass) must not suppress carry. + if not _ir_value_defined_inside_op(ref_val, loop_op): + return False + # The nested loop carries the leaf across loop_op iterations only when + # its ref re-inits from a value loop_op itself carries. + try: + ref_init = op.operands[0] + except Exception: + return True + if isinstance(ref_init, ir.Value) and _pyir_value_rooted_outside_loop( + ref_init, loop_op + ): + return False + return True + except Exception: + return False -def _raw_backing_ir_value(arg: Any) -> "ir.Value | None": - """Return *arg*'s backing ``ir.Value`` WITHOUT any side effects. - - This must never invoke the value's ``ir_value()`` accessor: on - several DSL value types (``ArithValue``, ``Vector``) ``ir_value`` is - a ``@dsl_user_op`` whose wrapper re-runs ``_pyir_auto_load_arg`` on - the receiver, and on ``Numeric`` it routes through ``.to(ir.Value)`` - which constructs a fresh ``@dsl_user_op`` ``ArithValue``. Calling it - from ``_value_dominates_current_ip`` -- itself reached *from* - that auto-load wrapper -- creates unbounded mutual recursion that - overflows the C stack (SIGABRT). Instead read the already-baked SSA - value structurally: - - - ``ArithValue`` / ``Vector`` subclass ``ir.Value`` directly. - - ``Numeric`` stores its SSA value in ``.value`` (``ir.Value`` once - baked; a Python primitive when still a meta constant). - - Returns ``None`` when no baked ``ir.Value`` is available; the caller - then conservatively treats the value as not dominating and re-loads, - which is always sound (an extra ``pyir.load``, never wrong data). - """ - if isinstance(arg, ir.Value): - return arg - inner = getattr(arg, "value", None) - if isinstance(inner, ir.Value): - return inner - return None +def _pyir_holder_walks_align( + arg: Any, snapshot: "list[tuple[Any, str, ir.Value]]" +) -> bool: + """True iff *arg*'s ir.Value-holder walk corresponds record-for-record + ((holder class, attr name) pairs, in walk order) with *snapshot* -- the + rebound binding is a same-shaped value tree whose leaf holders nest + BELOW the binding (e.g. a wrapper delegating to an interior impl + object), so positional leaf pairing is faithful even though the + snapshot's immediate holder class differs from the binding's own.""" + if not snapshot: + return False + try: + cur_records = _pyir_walk_ir_value_holders(arg) + except Exception: + return False + if len(cur_records) != len(snapshot): + return False + for snap_rec, cur_rec in zip(snapshot, cur_records): + if type(snap_rec[0]) is not type(cur_rec[0]) or snap_rec[1] != cur_rec[1]: + return False + return True + + +def _pyir_resolve_loop_leaf_updates( + arg: Any, + snapshot: "list[tuple[Any, str, ir.Value]]", +) -> "list[tuple[Any, str, ir.Value, ir.Value, Any, int]]": + """Pair each pre-loop snapshot leaf with the loop body's CURRENT leaf value.""" + updates: "list[tuple[Any, str, ir.Value, ir.Value, Any, int]]" = [] + # Whole-object-rebind support: snapshot owner's and current binding's canonical leaves, aligned + # positionally. Computed lazily (the common in-place path never pays this nor risks a mis-pair). + snap_owner = snapshot[0][0] if snapshot else None + snap_extract: "list[ir.Value] | None" = None + cur_extract: "list[ir.Value] | None" = None + rebind_aligned = False # snapshot/current extract lined up by length + + for snap_idx, (snap_holder, attr_name, snap_value) in enumerate(snapshot): + cur_value: Any = None + rebind_holder: Any = snap_holder + reconstruct_obj: Any = None + leaf_index = -1 + # 0. SELF-LEAF: the snapshot is a whole-object bare ``ir.Value`` (the local IS + # its own leaf). + if _is_self_leaf_record(snap_holder, attr_name): + # OPAQUE self-leaves only: a scalar bare ``ir.Value`` is already + # carried by the slot machinery; carrying it here would double-carry. + if not _is_opaque_leaf_type(snap_value): + continue + cur_self = _raw_backing_ir_value(arg) + if isinstance(cur_self, ir.Value) and cur_self is not snap_value: + cur_value = cur_self + if cur_value is None: + continue + updates.append( + ( + rebind_holder, + attr_name, + snap_value, + cur_value, + reconstruct_obj, + leaf_index, + ) + ) + continue + # 1. In-place attribute mutation on the original holder. + try: + in_place = getattr(snap_holder, attr_name) + except AttributeError: + in_place = None + if isinstance(in_place, ir.Value) and in_place is not snap_value: + cur_value = in_place + rebind_holder = snap_holder + elif arg is not None and arg is not snap_owner: + # 2. Whole-object rebind (local binds a NEW object): align this snapshot leaf to a canonical + # extract leaf of the owner, then to the same position in the current binding's extract. + if snap_extract is None: + # The owner is the rebound object itself (same class as the + # current binding), or the snapshot's holders nest BELOW the + # binding as a same-shaped value tree (a wrapper delegating + # to an interior impl -- e.g. Array -> _ArrayImpl): both give + # a faithful positional leaf pairing. + snap_extract = _pyir_extract_any_leaf_values(snap_owner) or [] + cur_extract = _pyir_extract_any_leaf_values(arg) or [] + rebind_aligned = ( + len(snap_extract) > 0 + and len(snap_extract) == len(cur_extract) + and ( + type(snap_owner) is type(arg) + or _pyir_holder_walks_align(arg, snapshot) + ) + ) + if rebind_aligned and snap_extract is not None and cur_extract is not None: + # Participates only if this leaf's SSA is one of the owner's canonical extract + # leaves; a derived holder is skipped. ``_same_ir_value`` (same underlying MLIR + # value), not Python ``is`` -- extraction may re-wrap the same SSA in a fresh + # ``ir.Value`` object -- and not ``==`` (structural). + matched = [ + _k + for _k, _ev in enumerate(snap_extract) + if _same_ir_value(_ev, snap_value) + ] + ci = matched[0] if matched else -1 + if ci >= 0 and _is_opaque_leaf_type(snap_value): + # DUPLICATE SSA among the owner's extract leaves: every record + # collapses onto the FIRST match, so a duplicated position's own + # update would silently cross-bind. Refuse unless every + # duplicated position is unchanged (then no pairing is needed). + if len(matched) > 1 and any( + not _same_ir_value(cur_extract[_k], snap_value) + for _k in matched + ): + _pyir_raise_rebind_duplicate_leaf(arg) + # OPAQUE leaves only: a scalar rebind is already carried by the M2S function-entry ref, + # so carrying it again would double-carry. + new_leaf = cur_extract[ci] + if isinstance(new_leaf, ir.Value) and not _same_ir_value( + new_leaf, snap_value + ): + cur_value = new_leaf + rebind_holder = snap_holder + reconstruct_obj = arg + leaf_index = ci + if cur_value is None: + continue + updates.append( + ( + rebind_holder, + attr_name, + snap_value, + cur_value, + reconstruct_obj, + leaf_index, + ) + ) + return updates -def _value_dominates_current_ip(value: Any) -> bool: - """Return ``True`` if *value*'s backing ``ir.Value`` dominates the - current MLIR insertion point. - - Accepts a DSL-wrapped value, a bare ``ir.Value``, or ``None``; the backing - SSA is read via :func:`_raw_backing_ir_value` (not ``value.ir_value()``, - which would re-enter the auto-load path that calls this and recurse). - - Used by ``_pyir_auto_load_arg`` (reuse a load vs emit a fresh one) and the - region-escape repair (is a leaf trapped in a closed region). Conservative: - returns ``False`` on any error, no ``pyir``, or no baked SSA, so the caller - re-loads / repairs rather than keep a non-dominating value. - - Implemented in C++ (``pyir.value_dominates_ip``): ONE binding crossing - folds the accessibility leg (ancestor region, ``IsolatedFromAbove`` - barriers) and the same-block ordering leg (the ancestor-region check alone - is necessary but not sufficient -- a same-region value may be defined - after the IP, so a same-block def must precede the IP's reference - operation, and any same-block def counts when the IP is at the block - end). This probe runs per staged arg of every emitted op, so the - step-by-step Python walk it replaces was the highest-frequency - binding-crossing overhead of the trace. - """ +def _pyir_load_dominates_use(loaded: "ir.Value", user_op: "ir.Operation") -> bool: + """Return ``True`` if *loaded* (a body-entry ``pyir.load`` at a region block's begin) + DOMINATES *user_op*.""" if pyir is None: return False try: - raw = _raw_backing_ir_value(value) - if raw is None: + user_block = user_op.block + if user_block is None: return False - ip = ir.InsertionPoint.current - ref_op = ip.ref_operation - if ref_op is not None: - ref_op = getattr(ref_op, "operation", ref_op) - return bool(pyir.value_dominates_ip(raw, ip.block, ref_op)) + if not pyir.is_value_in_ancestor_region(loaded, user_block): + return False + load_op = _get_defining_operation(loaded) + if load_op.block != user_block: + # User is in a region strictly nested below the load's block -> dominated. + return True + # Same block: the load sits at block begin, so it dominates any op after it; + # ``is_before_in_block`` orders the two within the block. + return load_op.is_before_in_block(getattr(user_op, "operation", user_op)) except Exception: return False -def _same_ir_value(a: "ir.Value | None", b: "ir.Value | None") -> bool: - """Whether *a* and *b* are the SAME SSA value, WITHOUT emitting IR. - - Two distinct Python wrappers can refer to one underlying MLIR value (the - bindings re-mint a wrapper per access), so ``id()`` is unreliable; but a DSL - ``ArithValue`` overrides ``==`` to emit an ``arith.cmpi`` (an op leaking into - the IR at the current insertion point -- a dominance hazard). Compare via the - BASE ``ir.Value.__eq__`` (a nanobind pointer comparison on the C++ value), - bypassing any subclass override. Conservative: any error -> ``False``. - """ - if a is None or b is None: +def _pyir_typed_update_is_advance( + cur_value: "ir.Value", + sources: "list[ir.Value]", + region_op: "ir.Operation", + max_ops: int = 512, +) -> bool: + """Whether *cur_value*'s def chain inside *region_op* transitively consumes any of + *sources* -- the old leaf value (or its in-body loads).""" + if not sources: return False try: - return bool(ir.Value.__eq__(a, b)) + seen: "set[int]" = set() + # Pin every visited wrapper for the walk's duration: a GC'd wrapper's + # recycled id would alias a NEW op into ``seen`` and break the chain. + pins: "list[Any]" = [] + work: "list[ir.Value]" = [cur_value] + budget = max_ops + while work: + if budget <= 0: + # Exhaustion is UNKNOWN, never a silent skip: report an + # advance so the caller refuses loudly (bounds gate by loud + # refusal, not truncation). + return True + v = work.pop() + for src in sources: + if _same_ir_value(v, src): + return True + owner = v.owner + if isinstance(owner, ir.Block): + continue # block argument: no defining op to walk through + def_op = getattr(owner, "operation", owner) + oid = id(def_op) + if oid in seen: + continue + seen.add(oid) + pins.append(def_op) + if not _op_is_inside_op(def_op, region_op): + continue # chain exits the region: loop-invariant root + budget -= 1 + for operand in def_op.operands: + work.append(operand) except Exception: - return a is b - + return False + return False -def _op_is_inside_op(op: "ir.Operation", outer_op: "ir.Operation") -> bool: - """Return ``True`` if *op* is nested (at any depth) inside *outer_op*. - Walks *op*'s ``parent`` chain (op -> enclosing op) comparing each ancestor - to *outer_op*. Uses only ``Operation.parent`` (never ``Block.region`` / - ``Block.owner``, whose Python bindings can fault on certain blocks). - Conservative: any error returns ``False``. - """ +def _pyir_ref_use_inside_op( + ref: "ir.Value", loop_op: "ir.Operation", op_name: str +) -> bool: + """Return True if a use of *ref* by an *op_name* op exists inside *loop_op*.""" + if pyir is None: + return False try: - target = getattr(outer_op, "operation", outer_op) - cur = getattr(op, "operation", op) - while cur is not None: - if cur == target: + for use in ref.uses: + user = getattr(use, "owner", None) + if user is None: + continue + user_op = getattr(user, "operation", user) + try: + if str(user_op.name) != op_name: + continue + except Exception: + continue + if _op_is_inside_op(user_op, loop_op): return True - parent = cur.parent - if parent is None: - return False - cur = getattr(parent, "operation", parent) - return False except Exception: return False + return False -def _innermost_enclosing_loop_op_at_ip() -> "ir.Operation | None": - """Return the innermost ``scf.for`` / ``scf.while`` op enclosing the - current insertion point, or ``None`` if there is no enclosing loop. - - Mirrors :func:`_get_function_entry_block`'s upward walk: it ascends from - ``InsertionPoint.current`` using only the ``Block.owner`` (block -> op) and - ``Operation.block`` (op -> parent block) navigation that is proven not to - fault on the binding's recycled CF blocks, stops at the nearest function / - module boundary (never walking into a recycled module-block wrapper), and - bounds the walk by IR depth rather than an ``id()`` visited-set. Unlike - :func:`_op_has_enclosing_loop` (which answers a yes/no question about a - given op), this returns the loop *op* so callers can pass it to the - ``Operation.parent``-based containment tests (:func:`_ir_value_defined_inside_op`). - Conservative: any error -> ``None`` (caller treats the carrying loop as - indeterminate). - """ - try: - block = ir.InsertionPoint.current.block - except Exception: - return None +def _pyir_slot_stored_in_body(ref: "ir.Value", loop_op: "ir.Operation") -> bool: + """Return True if a ``pyir.store`` into *ref* exists inside *loop_op*.""" + return _pyir_ref_use_inside_op(ref, loop_op, "pyir.store") - _MAX_NESTING = 256 - for _ in range(_MAX_NESTING): - if block is None: - return None - try: - parent = block.owner # Block -> Python dialect op - except Exception: - return None - op = getattr(parent, "operation", parent) # -> ir.Operation - try: - op_name = str(op.name) - except Exception: - return None - if op_name in ("scf.for", "scf.while"): - return op - # A function / module boundary owns no enclosing loop; stop rather than - # walk into recycled top-of-region block aliasing (a SIGSEGV hazard). - if _is_func_boundary_op(op_name) or _is_module_boundary_op(op_name): - return None - try: - block = op.block # ir.Operation -> parent Block - except Exception: - return None + +def _pyir_ref_loaded_inside_op(ref: "ir.Value", loop_op: "ir.Operation") -> bool: + """Return True if a ``pyir.load`` of *ref* exists inside *loop_op*.""" + return _pyir_ref_use_inside_op(ref, loop_op, "pyir.load") + + +def _pyir_unwrap_meta_primitive(value: Any) -> "Any | None": + """Return the Python primitive a meta scalar wraps, else ``None``.""" + # Wrapper unwrap precedes the bare-primitive test: a wrapper subclasses + # int/float, so isinstance would return it (a wrapped bool then mints i32). + pv = getattr(value, "python_value", None) + if type(pv) in (bool, int, float): + return pv + if type(value) in (bool, int, float): + return value + inner = getattr(value, "value", None) + if isinstance(inner, (bool, int, float)): + return inner return None -def _loop_free_enclosing_if_ops_at_ip() -> "list[ir.Operation]": - """Return every ``scf.if`` op enclosing the current insertion point, from - innermost outward, stopping at the FIRST ``scf.for`` / ``scf.while`` (or a - function / module boundary). - - Mirrors :func:`_innermost_enclosing_loop_op_at_ip`'s upward walk, using - only the ``Block.owner`` / ``Operation.block`` navigation proven not to - fault on recycled CF blocks. The returned list is the chain of standalone - (loop-free) ``scf.if`` ancestors -- a caller checks each for a sibling-branch - re-bake, since the relevant ``if`` (the mega-kernel ``if pref: ... else: - ...``) may be an OUTER ancestor while an inner ``scf.if`` sits between it - and the IP. An empty list means no loop-free enclosing ``scf.if``. - Conservative: any error -> the chain collected so far. - """ - ifs: "list[ir.Operation]" = [] - try: - block = ir.InsertionPoint.current.block - except Exception: - return ifs +def _is_region_meta_scalar(value: Any) -> bool: + """Return ``True`` for a value-tree field that is a plain *meta* scalar.""" + if isinstance(value, ir.Value): + return False + if value is None or isinstance(value, (bool, int, float, str, bytes)): + return True + return isinstance(value, _WatchedM) - _MAX_NESTING = 256 - for _ in range(_MAX_NESTING): - if block is None: - return ifs + +def _pyir_walk_meta_holders(obj: Any) -> "list[tuple[Any, str, Any]]": + """Return the in-place holders of every meta-scalar leaf in *obj*.""" + return _walk_value_tree_holders(obj, _is_region_meta_scalar) + + +def _pyir_restore_region_meta( + arg: Any, + snapshot: "list[tuple[Any, str, Any]] | None", + context: str, +) -> Any: + """Restore value-tree meta-scalar leaves mutated inside a staged region.""" + if snapshot is None: + return arg + for holder, attr_name, snap_value in snapshot: try: - parent = block.owner - except Exception: - return ifs - op = getattr(parent, "operation", parent) + cur_value = getattr(holder, attr_name) + except AttributeError: + continue + # Identity short-circuit covers the no-mutation and interned-small-int cases. Only rebind when the + # body left another meta scalar there (else the value is not ours to restore). + if cur_value is snap_value or not _is_region_meta_scalar(cur_value): + continue try: - op_name = str(op.name) - except Exception: - return ifs - if op_name in ("scf.for", "scf.while"): - return ifs - if op_name == "scf.if": - ifs.append(op) - if _is_func_boundary_op(op_name) or _is_module_boundary_op(op_name): - return ifs + _pyir_setattr_raw(holder, attr_name, snap_value) + except (AttributeError, TypeError): + continue + # Forget the closed region's baked-constant bookkeeping so a sibling region + # re-reading the leaf starts from op-entry state, not spuriously promoted. + _forget_region_local_slot_state(holder, attr_name) + log().info( + "[pyir region] restored region-mutated meta leaf '%s' (%s)", + attr_name, + context, + ) + return arg + + +def _walk_meta_numeric_leaf_holders(obj: Any) -> "list[tuple[Any, str, Any]]": + """Return parent-attribute holders of every META-backed Numeric leaf in *obj*.""" + return _walk_value_tree_holders(obj, _is_meta_numeric_leaf) + + +def _is_meta_numeric_leaf(value: object) -> bool: + """Return True if *value* is a staged scalar Numeric WRAPPER whose backing is still a + Python primitive (a meta literal, not a baked SSA).""" + return _numeric_leaf_kind(value) == _NUMERIC_LEAF_META + + +def _pyir_restore_meta_numeric_leaves( + snapshot: "list[tuple[Any, str, Any]] | None", + context: str, +) -> None: + """Restore a write-only literal-origin Numeric leaf the closed region baked. + Companion to :func:`_walk_meta_numeric_leaf_holders`.""" + if snapshot is None: + return + for holder, attr_name, snap_value in snapshot: try: - block = op.block - except Exception: - return ifs - return ifs + cur_value = getattr(holder, attr_name) + except AttributeError: + continue + if cur_value is snap_value: + continue + cur_raw = _raw_backing_ir_value(cur_value) + if cur_raw is None: + # Still a meta literal -- nothing baked / trapped. + continue + if _value_dominates_current_ip(cur_raw): + continue + if _pyir_leaf_is_carried(cur_value): + # A real loop-carried accumulator carried through a slot is authoritative; + # only an untracked dead write-only counter is restored. + continue + try: + _pyir_setattr_raw(holder, attr_name, snap_value) + except (AttributeError, TypeError): + continue + log().info( + "[pyir region] restored region-baked write-only meta numeric leaf " + "'%s' (%s)", + attr_name, + context, + ) + + +def _is_untracked_post_region_binding(value: Any) -> bool: + """Return True if *value* is a post-region local binding that the + ``_pyir_auto_load_arg`` merge could NOT carry through a ref.""" + try: + if isinstance(value, _WatchedM): + return False + if _get_load_version(value) is not None: + return False + mv = getattr(value, "_mutable_ref", None) + if mv is not None and mv._is_ref_accessible(): + return False + return True + except Exception: + return False -__all__ = [name for name in list(globals()) if not name.startswith("__")] +# Static export surface (regenerate with scripts/gen_pyir_all.py): the +# module's OWN names -- its namespace minus what lower chain modules +# already export and minus the generator's INTERNAL_ONLY registry. +__all__ = [ + "_functools", + "gc", + "_sys", + "weakref", + "NoReturn", + "_PYIR_LOAD_VERSION_ATTR", + "_PYIR_BOUNDARY_META_CELL_READS", + "_PYIR_STRUCTURAL_META_CONSUMPTIONS", + "_PYIR_STRUCTURAL_META_CONSUMPTION_SITES", + "_PYIR_ARM_LOCAL_META_WRITES", + "_PYIR_BOUNDARY_FLIP_GUARD_READS", + "_META_CONST_REPLACEMENTS", + "_ref_write_epoch", + "_get_load_version", + "_get_load_depth", + "_get_load_epoch", + "_pyir_adopt_stored_representative", + "_is_func_boundary_op", + "_is_module_boundary_op", + "_auto_promote_primitive", + "_can_create_ref", + "_is_scalar_ssa_carryable", + "_can_carry_leaf_ref", + "_is_vector_like", + "_mlir_types_match", + "_types_match", + "_is_memref_like", + "_pyir_raise_memref_inregion_alloc_rebind", + "_pyir_type_is_register_memref", + "_PYIR_MEMREF_LAST_STORE", + "_pyir_note_memref_store", + "_pyir_scratch_stale_serve_record", + "_pyir_refuse_stale_memref_serve", + "_pyir_classify_transitive_consumers", + "_pyir_load_in_rebind_reexecution_scope", + "_pyir_scratch_value_admitted", + "_pyir_ref_pointee_type_changed", + "_pyir_raise_type_changed_in_region", + "_pyir_raise_rebind_duplicate_leaf", + "_pyir_record_slot_template", + "_declared_m2s_promotion_class", + "_pyir_declared_promotion_template", + "_is_boolean_like", + "_is_literal_backed", + "_wrap_ir_like", + "_make_poison_like", + "_make_zero_like", + "_pyir_type_is_scalar", + "_make_raw_placeholder_init", + "_first_non_dsl_caller_location", + "_POISON_EMITTED", + "_get_defining_operation", + "_get_function_entry_block", + "_pyir_recorded_birth_block", + "_mint_region_born_cell", + "_create_ref", + "_pyir_emit_store", + "_SlotId", + "_NON_WEAKREF_OWNER_TYPES_SEEN", + "_make_slot_id", + "_cell_home_binding", + "_ce_local_key", + "_ce_note_local_assign", + "pyir_seed_param_bindings", + "_make_slot_key", + "_pyir_owner_slot_is_computed", + "_pyir_read_place", + "_pyir_route_is_live_place_row", + "_pyir_route_is_current", + "_pyir_resolve_snapshot", + "_pyir_row_binding_unobserved_write", + "_push_scope", + "_pop_scope", + "_PyirScopeGuard", + "pyir_register_scope_cells", + "pyir_register_nonlocal_names", + "_pyir_boundary_cell_read_for_local", + "_owner_token", + "_pyir_wall_composite_str_slot", + "_place_for", + "_corrected_place_for", + "_register_place_prefix", + "_pyir_same_container_structure", + "_pyir_propagate_owner_token", + "_pyir_adopt_rebuilt_owner_token", + "_pyir_lookup_owner_token", + "_pyir_validate_owner_class", + "_pyir_owner_is_celled", + "_pyir_emission_self_check", + "_pyir_guard_stale_epoch", + "_pyir_witness_predicate_fold", + "_pyir_check_staged_fold_witness", + "_pyir_module_is_tracer_layer", + "_pyir_record_external_payload_consumption", + "_pyir_record_staged_identity_hashed", + "_pyir_staged_hash_witness_count", + "_emit_constant_at_current_ip", + "_emit_constant_for_ref", + "_const_value_of", + "_const_values_equal", + "_unwrap", + "_pyir_structural_value_conflicts", + "_pyir_structural_bake_is_reseeded", + "_pyir_merge_src_pairs", + "_WatchedM", + "_WatchedInt", + "_WatchedBool", + "_WatchedFloat", + "_PYIR_DICT_CF_CREATED_KEYS", + "_WatchedDict", + "_WatchedList", + "_pyir_identity_compare_choke", + "_pyir_watched_dict_label", + "_pyir_watched_list_label", + "_pyir_adopt_dict_value", + "_pyir_instance_dict_owner", + "_pyir_spec_value_is_ir_wrapper", + "_pyir_spec_chain_value", + "_pyir_spec_stamp_trace_born", + "_pyir_adopt_list_value", + "_pyir_register_container_held_object", + "_pyir_owner_is_container_held", + "_SEQ_WALK_INERT_PRIMS", + "_mark_container_walk_visited", + "_sequence_walk_is_inert", + "_pyir_adopt_containers_under", + "_pyir_record_host_restore", + "_pyir_holder_store", + "_replace_value_uses", + "_exit_function_trace", + "_pyir_take_sealed_spec", + "_pyir_assert_entry_attested", + "_slot_storage_available", + "_get_slot_mv", + "_pyir_track_slot_holder", + "_pyir_refuse_del_finalizer_owner", + "_pyir_sighting_inert_class", + "_pyir_register_candidate_holder", + "_pyir_settle_sighting", + "_pyir_register_trace_args", + "_pyir_spec_signature_facts", + "_pyir_spec_seg_steps", + "_pyir_spec_observe_attr_read", + "_pyir_spec_observe_global_read", + "_pyir_spec_record_read", + "_pyir_spec_record_attr_probe", + "_pyir_spec_note_fabricated_bake", + "_pyir_spec_record_write", + "_pyir_spec_record_unbind", + "_pyir_import_is_cached", + "_pyir_spec_record_import", + "_pyir_spec_structural_amend", + "_pyir_stamp_rewritten", + "_pyir_callee_is_rewritten", + "_pyir_spec_boundary_closure_read", + "_pyir_boundary_taken_defaults", + "_pyir_register_taken_default_holders", + "_pyir_spec_boundary_default_reads", + "_pyir_registry_candidate_objects", + "_set_slot_mv", + "_iter_owner_slot_mvs", + "_PyirGenEventSuppress", + "_pyir_gen_events_suppressed", + "_pyir_is_generation_compound", + "_pyir_keepalive_generation_obj", + "_pyir_record_compound_binding", + "_pyir_record_rebind_ref_cell", + "_pyir_record_generation_rebind", + "_pyir_complete_generation_rebind_cells", + "_pyir_raise_superseded", + "_pyir_check_superseded_owner_write", + "_pyir_check_superseded_wrapper_load", + "_pyir_generation_read_checks", + "_pyir_record_cf_attr_first_def", + "_pyir_update_cf_attr_first_def_on_write", + "_pyir_check_cf_attr_first_def_read", + "_attach_mutable_ref", + "_pyir_adopt_live_place_cell", + "_pyir_place_is_index_sibling", + "_pyir_route_restructured_tuple_leaves", + "_pyir_refuse_superseded_row_serve", + "_fresh_wrapper", + "_same_ir_value", + "_raw_backing_ir_value", + "_as_arith_capable_scalar_leaf", + "_value_dominates_current_ip", + "_op_is_inside_op", + "_block_strictly_inside", + "_block_inside_op", + "_pyir_write_in_if_arms_of", + "_pyir_check_arm_local_escape", + "_ir_value_defined_inside_op", + "_op_has_enclosing_loop", + "_innermost_enclosing_loop_op_at_ip", + "_loop_free_enclosing_if_ops_at_ip", + "_meta_use_in_sibling_if_region", + "_safe_instance_dict", + "_slots_member_descriptors", + "_has_instance_storage", + "_instance_storage_items", + "MutableValue", + "_ref_dominates_whole_function", + "_get_instance_attrs", + "_implements_dynamic_expression", + "_rebuild_tuple_like", + "_is_compound_single_leaf", + "_check_all_fields_decomposable", + "_flatten_tuple", + "_pyir_assign_simple", + "_pyir_record_fresh_object_leaf_first_defs", + "_subobject_is_ctor_private", + "_record_fresh_leaf_first_defs_walk", + "_record_meta_primitive_first_def", + "_mlir_type_or_none", + "_staged_type_changed", + "_pyir_lookup_slot_from_value", + "_pyir_value_tracked_by_accessible_ref", + "_pyir_row_load_type_mismatch", + "_load_as_dsl", + "_clear_slot_mv", + "_pyir_retire_place_row", + "_pyir_region_fresh_raw", + "_pyir_refresh_cell_read", + "_pyir_recovery_holder_pairs", + "_pyir_meta_primitive_value", + "_pyir_snapshot_registry_meta_slots", + "_pyir_verify_registry_meta_slots", + "_PYIR_OPAQUE_OWNER_AUDITS", + "_pyir_opaque_owner_key_sets", + "_pyir_arm_opaque_owner_audit", + "_pyir_verify_opaque_owner_audits", + "_ref_lives_in_strictly_enclosing_block", + "_carries_foreign_slot_binding", + "_freeze_foreign_slot_binding", + "_pyir_same_ref_reload", + "_is_opaque_leaf_type", + "_pyir_ir_type_is_opaque", + "_pyir_extract_leaf_values", + "_pyir_is_carryable_tuple", + "_pyir_extract_any_leaf_values", + "_pyir_new_from_mlir_values_any", + "_is_opaque_leaf_value_tree", + "_value_trees_different_prototype", + "_has_accessible_loop_carried_ref", + "_pyir_skip_trapped_child_region_writeback", + "_pyir_walk_ir_value_holders", + "_walk_ir_value_holders_into", + "_pyir_build_gather_segment", + "_is_self_leaf_record", + "_self_leaf_snapshot", + "_is_numeric_leaf_holder", + "_pyir_record_promoted_place_leaf", + "_pyir_reload_stale_staged_attr_leaves", + "_pyir_repair_captured_escaped_leaves", + "_pyir_repair_region_escaped_leaves", + "_pyir_repair_region_escaped_tuple_leaves", + "_pyir_reload_trapped_tuple_leaf", + "_pyir_local_place_for_name", + "_pyir_region_carry_place_for", + "_pyir_adopt_region_carry_ref", + "_pyir_adopt_d1_carry_cell", + "_pyir_single_opaque_extract_leaf", + "_pyir_pair_place_cell_handle", + "_pyir_place_opaque_rebind_choke", + "_pyir_memref_alloc_rooted_inside_region", + "_pyir_value_is_inner_loop_carried", + "_pyir_holder_walks_align", + "_pyir_resolve_loop_leaf_updates", + "_pyir_load_dominates_use", + "_pyir_typed_update_is_advance", + "_pyir_slot_stored_in_body", + "_pyir_ref_loaded_inside_op", + "_pyir_unwrap_meta_primitive", + "_pyir_walk_meta_holders", + "_pyir_restore_region_meta", + "_walk_meta_numeric_leaf_holders", + "_pyir_restore_meta_numeric_leaves", + "_is_untracked_post_region_binding", +] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_corewalk.py b/python/CuTeDSL/cutlass/base_dsl/pyir_corewalk.py index 52b25ae7a1..33d0e2b102 100644 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_corewalk.py +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_corewalk.py @@ -12,183 +12,135 @@ """PyIR runtime -- corewalk layer; see facade for the public surface.""" -from .pyir_core import * # noqa: F401,F403 (re-export lower layers up the chain) - -# -- BEGIN explicit imports for the type checker (do not edit the list by hand; -# it mirrors names the chain re-exports at runtime via the wildcard + dynamic -# ``__all__`` above, which a static type checker cannot evaluate -- so every -# name is also imported explicitly from the layer that DEFINES it). Purely -# additive: the wildcard import stays the runtime source of truth. -from .pyir_state import ( # noqa: F401 - Any, - _is_staged_value, - _slot_refs, - ir, - pyir, -) -from .pyir_core import ( # noqa: F401 - _WatchedM, - _can_create_ref, - _check_all_fields_decomposable, - _create_ref, - _flatten_tuple, - _get_instance_attrs, - _load_as_dsl, -) -# -- END explicit imports for the type checker - - -def _pyir_auto_load_arg(arg: Any) -> Any: +from .pyir_state import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_core import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) + + +def _pyir_auto_load_arg(arg: Any, *, row_authoritative: bool = False) -> Any: """If *arg* carries a ``_mutable_ref`` with an accessible ref, emit ``pyir.load`` and return a fresh value that dominates the current insertion point. Otherwise return *arg* unchanged. - Called by ``@dsl_user_op`` on every positional argument so that - post-loop uses of pyir-tracked variables automatically get a load - from the ref instead of using a stale SSA value from inside the - loop body. - - Optimization: skip the auto-load when *arg* was itself produced by - the most recent ``MutableValue.load()`` (same ``_load_version``) - AND its underlying ``ir.Value`` still dominates the current - insertion point. This eliminates the redundant load that would - otherwise follow an AST-inserted ``pyir_read``. A ``store()`` - bumps ``_load_version`` and invalidates the cache. - - D1: when *arg* is a ``_WatchedM`` wrapper: - * If the slot has been promoted to ``pyir.ref`` (an entry in - ``_slot_refs``), emit ``pyir.load`` and wrap in the matching - DSL Numeric. - * Otherwise return the wrapper unchanged. ``_WatchedInt`` / - ``_WatchedBool`` / ``_WatchedFloat`` subclass ``int`` / ``float`` - so consumers that use the arg as a Python primitive (e.g. - ``ir.VectorType.get([vec_size], ...)`` reading ``vec_size`` as - the shape, ``isinstance(x, int)`` branches, ``const_expr``) - work transparently. Consumers that need an SSA value - (Numeric ``__init__``, ``arith.const``, ``cute.printf`` arg - coercion) already invoke ``.ir_value()`` on the wrapper, which - performs the bake AND records into ``_meta_uses`` so D1's - retroactive rewrite still fires when the slot later mutates. + Called by ``@dsl_user_op`` on every positional argument so post-loop + uses of pyir-tracked variables load from the ref instead of using a + stale SSA value from inside the loop body. That user channel follows + the route only while *arg* is provably the cell's current binding + (V-2); a retained snapshot keeps its own SSA (LAW 1). + + *row_authoritative* marks the carry engine's post-region rebind of + a carried binding: there the region just wrote the row, the row IS the + authority, and the reload deliberately skips the snapshot judgment + (not the superseded-scratch refusal: a sound carry presents the cell's + current handle, so a superseded handle here is never legitimate). + + Optimization: skip the auto-load when *arg* is itself the most recent + ``MutableValue.load()`` (same ``_load_version``) AND still dominates the + current IP -- eliminating the redundant load after an AST ``pyir_read``. + A ``store()`` bumps ``_load_version`` and invalidates the cache. + + D1, when *arg* is a ``_WatchedM`` wrapper: if the slot was promoted to a + ref (in ``_slot_refs``), emit ``pyir.load`` wrapped in the matching DSL + Numeric; otherwise return the wrapper unchanged. ``_WatchedM`` subclasses + int/float so primitive consumers (shapes, ``isinstance``, ``const_expr``) + work transparently, while SSA consumers invoke ``.ir_value()`` (which + bakes AND records into ``_meta_uses`` so retroactive rewrite still fires). """ + # Candidate-holder registry: every ``@dsl_user_op`` argument is a sighting + # of an object the trace touches; the recursion below registers elements. + _pyir_register_candidate_holder(arg) + if isinstance(arg, _WatchedM): slot = arg._slot_key if slot is not None and slot in _slot_refs: - return _load_as_dsl(_slot_refs[slot], arg.python_value) - # Unpromoted: return the bare wrapper so the consumer decides - # whether to bake (via ``.ir_value()``) or use as a Python - # primitive (via the int/float subclass nature of ``_WatchedM``). + _ref = _slot_refs[slot] + # One declared SNAPSHOT semantics: the cell serves this wrapper's + # read only while it still holds the creation-time value (same + # ref, same write epoch); a superseded wrapper materializes + # through ir_value()'s snapshot judgment instead. + if _ref is getattr(arg, "_pyir_birth_ref", None) and _ref_write_epoch( + _ref + ) == getattr(arg, "_pyir_birth_epoch", None): + return _load_as_dsl(_ref, place=slot, stamp_place=True) + return arg + # Unpromoted: return the bare wrapper so the consumer decides whether to + # bake (``.ir_value()``) or use it as a Python primitive. return arg - # Tuple recursion: tuples have no ``_mutable_ref`` themselves but - # their staged leaves often do. Mirrors ``pyir_read``'s existing - # tuple recursion so ``@dsl_user_op`` boundaries like - # ``Vector.from_elements((rm, rm, ...))`` auto-load each element. - if isinstance(arg, tuple): - return tuple(_pyir_auto_load_arg(e) for e in arg) + # Tuple/list recursion: containers have no ``_mutable_ref`` but their staged + # leaves often do -- auto-load each element to mirror ``pyir_read``. + if isinstance(arg, (tuple, list)): + loaded = [ + _pyir_auto_load_arg(e, row_authoritative=row_authoritative) for e in arg + ] + # ``type(arg)(genexpr)`` collapses / breaks a namedtuple; rebuild via + # the tuple-aware primitive and keep a list a list. + return loaded if isinstance(arg, list) else _rebuild_tuple_like(arg, loaded) + + # Owner-less retained-handle read (e.g. a container element's leaf) reaches + # this auto-load with a superseded generation's own leaf wrapper. + if _SUPERSEDED_LEAF_WRAPPERS: + _pyir_check_superseded_wrapper_load(arg) mv = getattr(arg, "_mutable_ref", None) if mv is not None and mv.ref is not None: if mv._is_ref_accessible(): # Dedup: arg is the latest load AND still dominates current IP. - cached_version = getattr(arg, "_pyir_load_version", None) + cached_version = _get_load_version(arg) if ( cached_version is not None and cached_version == mv._load_version - and _arg_value_dominates_current_ip(arg) + # Depth-aware: only reuse a cached load at the SAME staged-CF depth, + # so a deeper consumer rebinds the inner loop's iter_arg. + and _get_load_depth(arg) == current_staged_cf_depth() + # Epoch-aware: a raw slot-ref store bumps the REF's write-epoch + # without touching this wrapper's version -- the load is stale. + and _get_load_epoch(arg) == _ref_write_epoch(mv.ref) + and _value_dominates_current_ip(arg) ): return arg - return mv.load() - # Ref exists but is inaccessible (stale, from sibling/exited CF). - # Re-create at current scope to restore dominance -- but only - # when ``arg``'s SSA value still dominates the current IP. If - # it doesn't, ``_create_ref`` would D-fallback to a - # ``ub.poison``-initialised ref at function entry with no - # accompanying store, and the subsequent ``mv.load()`` would - # silently read poison. In that case fall through to - # returning ``arg`` unchanged -- the verifier (or the - # end-of-trace poison-read catcher, if a sibling code path - # did create the same poison-init pattern) will surface the - # real issue at trace time instead of producing wrong results. - if _can_create_ref(arg) and _arg_value_dominates_current_ip(arg): + # An ADMITTED rebound memref cell must not re-serve a SUPERSEDED + # wrapper (a handle retained from before the last in-region + # rebind, e.g. through a container): every in-region allocation + # aliases the one entry-hoisted buffer, so the reload would + # observe the fresh generation's overwrites instead of the + # buffer Python retained. Refuse loudly at the use site. + _pyir_refuse_stale_memref_serve(arg, mv) + # R3: reload through the route when it IS its place's live row + # (place authority) or while *arg* is provably the cell's current + # binding (V-2); a retained snapshot keeps its own SSA (LAW 1). + if row_authoritative or _pyir_route_is_current(arg, mv): + return mv.load() + if _pyir_route_is_live_place_row(mv): + # LAW-1 counter-evidence: an epoch-tagged product a later + # store superseded, whose own SSA dominates, is a retained + # SNAPSHOT -- place authority does not govern that read. + _epoch_tag = _get_load_epoch(arg) + if ( + _epoch_tag is not None + and _epoch_tag != _ref_write_epoch(mv.ref) + and _value_dominates_current_ip(arg) + ): + return arg + return mv.load() + return _pyir_resolve_snapshot(arg, type(arg).__name__) + # Ref inaccessible (stale, from sibling/exited CF): re-create at current scope to restore + # dominance, but only when ``arg``'s SSA still dominates the IP -- else return ``arg``. + if _can_create_ref(arg) and _value_dominates_current_ip(arg): + # The guards above screen out the legitimate not-creatable cases, so + # a failure here means re-creation was expected to succeed but did not. try: new_mv = _create_ref(arg) return new_mv.load() - except Exception: - pass # fall through to return arg + except Exception as e: + raise DSLRuntimeError( + f"failed to re-create a dominating ref for a " + f"{type(arg).__name__} value after its original ref became " + f"inaccessible" + ) from e return arg -def _raw_ir_value(arg: Any) -> "ir.Value": - """Return *arg*'s backing ``ir.Value`` WITHOUT re-entering the - ``@dsl_user_op`` instrumentation in :func:`_pyir_auto_load_arg`. - - ``_arg_value_dominates_current_ip`` needs only the raw MLIR value to - inspect its owner/defining-op; it must not trigger op emission or a - nested auto-load. Calling ``arg.ir_value()`` directly is unsafe: - for DSL value types whose ``ir_value`` is itself ``@dsl_user_op`` - -wrapped (e.g. ``cute.TensorSSA``), the wrapper auto-loads its own - ``self`` argument, which calls back into ``_pyir_auto_load_arg`` -> - ``_arg_value_dominates_current_ip`` -> ``ir_value()`` -> ... ad - infinitum (RecursionError). - - DSL value wrappers (``ArithValue`` / ``Vector`` / ``TensorSSA``) - subclass ``ir.Value`` directly, so when *arg* is already an - ``ir.Value`` it IS its own raw MLIR value and no call is needed. - Otherwise fall back to the un-instrumented underlying ``ir_value`` - (``__wrapped__`` strips the ``@dsl_user_op`` wrapper) so the - dominance probe stays side-effect-free. - """ - if isinstance(arg, ir.Value): - return arg - iv = getattr(type(arg), "ir_value", None) - unwrapped = getattr(iv, "__wrapped__", None) - if unwrapped is not None: - return unwrapped(arg) - return arg.ir_value() - - -def _arg_value_dominates_current_ip(arg: Any) -> bool: - """Return ``True`` if *arg*'s backing ``ir.Value`` dominates the - current MLIR insertion point. - - Used by ``_pyir_auto_load_arg`` to decide whether a previously - loaded value can be reused without emitting a fresh - ``pyir.load``. Conservative: returns ``False`` on any error so - the caller falls back to re-loading. - - The region check via ``is_value_in_ancestor_region`` is necessary - but not sufficient: same-region values may still be defined after - the current insertion point. For same-block defs we additionally - require the defining op to precede the IP using - ``is_before_in_block`` against the IP's reference operation (or - accept any same-block op when the IP is at the block end). - """ - if pyir is None: - return False - try: - raw = _raw_ir_value(arg) - ip = ir.InsertionPoint.current - current_block = ip.block - if not pyir.is_value_in_ancestor_region(raw, current_block): - return False - owner = raw.owner - if isinstance(owner, ir.Block): - return True # block argument — dominates everything in its region - def_op = getattr(owner, "operation", owner) - if def_op.block != current_block: - return True # proper-ancestor region — structural dominance - ref_op = ip.ref_operation - if ref_op is None: - return True # IP at block end; def was emitted earlier in trace - ref_op = getattr(ref_op, "operation", ref_op) - if def_op == ref_op: - return False - return def_op.is_before_in_block(ref_op) - except Exception: - return False - - def _has_decomposable_staged_fields(obj: object) -> bool: """Return True if *obj* can be auto-decomposed into per-field ``pyir_assign`` calls. @@ -201,28 +153,182 @@ def _has_decomposable_staged_fields(obj: object) -> bool: return _check_all_fields_decomposable(obj, _visited=set()) -def _has_any_staged_content(obj: object) -> bool: +def _has_any_staged_content(obj: object, _visited: "set[int] | None" = None) -> bool: """Return True if any field (deeply) is a staged value. Used for the error vs passthrough decision: if the object has staged content that cannot be decomposed, we raise an error. If it has NO staged content, it is a pure meta replacement (harmless). + + A tuple/list passed directly (not as an object attribute) has no + ``__dict__``, so the attribute walk below would report it as having no + staged content. Check its leaves explicitly so a top-level sequence + carrying staged values is correctly recognised as loop-carried state + (mirrors the nested-sequence branch inside the attribute walk). + + *_visited* breaks self-referential object graphs: an object reachable from + itself (directly or via a cycle) is examined once, so the recursion cannot + diverge. Mirrors the cycle guard the sibling walk + (:func:`_check_all_fields_decomposable`) already uses. """ + if _visited is None: + _visited = set() + oid = id(obj) + if oid in _visited: + return False + _visited.add(oid) + + if isinstance(obj, (tuple, list)): + return any(_is_staged_value(e) for e in _flatten_tuple(tuple(obj))) for attr_name in _get_instance_attrs(obj): value = getattr(obj, attr_name) if _is_staged_value(value): return True - if isinstance(value, tuple) and any( - _is_staged_value(e) for e in _flatten_tuple(value) + if isinstance(value, (tuple, list)) and any( + _is_staged_value(e) for e in _flatten_tuple(tuple(value)) ): return True if ( - hasattr(value, "__dict__") + _has_instance_storage(value) and not isinstance(value, (int, float, bool, str, bytes, type)) - and _has_any_staged_content(value) + and _has_any_staged_content(value, _visited) ): return True return False -__all__ = [name for name in list(globals()) if not name.startswith("__")] +def _gather_captured_holders( + exclude_ids: "set[int] | None", + walk_obj: "Callable[[Any, set[int], list], Any]", +) -> "list[tuple[Any, str, Any]]": + """Shared registry sweep behind the captured-holder gathers. + + Collects leaf holders of objects CAPTURED by a staged region/loop body but + not passed as one of its ``mix_iter_args`` (captured free variables). + + Discovery iterates the trace-scoped candidate-holder registry + (:data:`_PYIR_CANDIDATE_HOLDERS`) in sighting order. + + Each live candidate is fed to *walk_obj*, which appends ``(holder, attr, + value)`` triples into the shared accumulator; the per-walker gate lives + entirely in *walk_obj*. + + One shared visited set spans all candidates, so a holder reachable from + several candidates is walked once and each ``(holder, attr)`` pair is + collected exactly once. + + Registration follows what the trace actually touched, not what happens to + sit in a Python frame, so a captured holder is reached at any stack depth. + + *exclude_ids* holds ``id()`` of objects already snapshotted via + ``mix_iter_args``. An empty result maps to the wrappers' sentinel. + + Cross-sweep caching (value-tree walk only): a cacheable root's record + segment (:func:`_pyir_build_gather_segment`) is memoized against its + write stamp and spliced at the root's registry position when the root + stamp AND every recorded Numeric leaf's own guard stamp are unchanged AND + no earlier live walk already consumed the root this sweep (positional + splice validation); any failing takes the live walk in position, so + record multiset AND order equal the uncached sweep's.""" + if exclude_ids is None: + exclude_ids = set() + holders: "list[tuple[Any, str, Any]]" = [] + visited_objs: "set[int]" = set() + if walk_obj is not _walk_ir_value_holders_into: + for cand in _pyir_registry_candidate_objects(): + if id(cand) in exclude_ids: + continue + walk_obj(cand, visited_objs, holders) + return holders + # Per-root decision flow (cached value-tree walk): + # + # root already visited this sweep? ──yes──► live walk (a no-op; replay + # │no here would emit duplicate rows) + # root has a write stamp? ──────────no───► live walk (no stamp to validate + # │yes a cached row against yet) + # cached row at the SAME stamp? + # │ ├─ negative row ────────────► live walk (judged uncacheable) + # │ ├─ a leaf guard stamp moved ► rebuild segment below + # │ └─ all guards intact ───────► replay cached (attr, val) rows + # │no (stale or missing row) + # build a fresh segment ── cacheable? ──yes──► store + replay it + # └──────────no──► store negative row + live walk + segments = _PYIR_GATHER_SEGMENTS + stamps = _PYIR_HOLDER_WRITE_STAMPS + for cand in _pyir_registry_candidate_objects(): + oid = id(cand) + if oid in exclude_ids: + continue + if oid in visited_objs: + # An earlier live walk consumed this root at ITS position. + walk_obj(cand, visited_objs, holders) + continue + stamp = stamps.get(oid) + if stamp is None: + walk_obj(cand, visited_objs, holders) + continue + row = segments.get(oid) + if row is not None and row[0] == stamp: + records = row[1] + if records is None: + # Negative row: judged not cacheable at this stamp; no + # re-judgment until the root is written again. + walk_obj(cand, visited_objs, holders) + continue + # Numeric-leaf guards: an in-place write on a recorded wrapper + # stamps the WRAPPER's row (not the root's) and may flip its leaf + # classification -- any moved stamp takes the rebuild instead. + for w, ws in row[2]: + if stamps.get(id(w)) != ws: + break + else: + visited_objs.add(oid) + for attr, val in records: + holders.append((cand, attr, val)) + continue + seg = _pyir_build_gather_segment(cand) + if seg is None: + segments[oid] = (stamp, None, ()) + walk_obj(cand, visited_objs, holders) + continue + records, guards = seg + segments[oid] = (stamp, records, guards) + visited_objs.add(oid) + for attr, val in records: + holders.append((cand, attr, val)) + return holders + + +class _PyirCapturedSnapshot(list): + """Captured-holder snapshot carrying the write clock read at gather START + (any write during or after the sweep stamps later, so the repair's + stamp<=clock skip is exact); behaves as a plain record list.""" + + __slots__ = ("pyir_gather_clock",) + pyir_gather_clock: int + + +def _pyir_gather_captured_leaf_holders( + exclude_ids: "set[int] | None" = None, +) -> "list[tuple[Any, str, ir.Value]] | None": + """Collect value-tree ``ir.Value`` leaf holders of objects CAPTURED by a staged + region body but not passed as one of its ``mix_iter_args``.""" + clock = _PYIR_WRITE_CLOCK[0] + holders = _gather_captured_holders(exclude_ids, _walk_ir_value_holders_into) + if not holders: + return None + snap = _PyirCapturedSnapshot(holders) + snap.pyir_gather_clock = clock + return snap + + +# Static export surface (regenerate with scripts/gen_pyir_all.py): the +# module's OWN names -- its namespace minus what lower chain modules +# already export and minus the generator's INTERNAL_ONLY registry. +__all__ = [ + "_pyir_auto_load_arg", + "_has_decomposable_staged_fields", + "_has_any_staged_content", + "_PyirCapturedSnapshot", + "_pyir_gather_captured_leaf_holders", +] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_entrypoints.py b/python/CuTeDSL/cutlass/base_dsl/pyir_entrypoints.py index 110a4ff2c3..ed1d009491 100644 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_entrypoints.py +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_entrypoints.py @@ -12,88 +12,25 @@ """PyIR runtime -- entrypoints layer; see facade for the public surface.""" -from .pyir_threading import * # noqa: F401,F403 (re-export lower layers up the chain) - -# -- BEGIN explicit imports for the type checker (do not edit the list by hand; -# it mirrors names the chain re-exports at runtime via the wildcard + dynamic -# ``__all__`` above, which a static type checker cannot evaluate -- so every -# name is also imported explicitly from the layer that DEFINES it). Purely -# additive: the wildcard import stays the runtime source of truth. -from .pyir_state import ( # noqa: F401 - Any, - DSLUserCodeError, - DiagId, - _PYIR_DICT_MUTATORS, - _PYIR_LIST_MUTATORS, - _PYIR_READ_SIMPLE_SENTINEL, - _PYIR_SET_MUTATORS, - _PYIR_SKIP, - _current_fn_id, - _is_staged_value, - _meta_uses, - _slot_first_def_block, - _slot_first_def_depth, - _slot_first_def_depth_any, - _slot_first_def_inside_cf, - _slot_mvs, - _slot_pending_store, - _slot_refs, - assign_meta_staged_check, - get_staged_cf_depth, - ir, - is_auto_m2s_enabled, - is_inside_constexpr_loop, - is_inside_locally_staged_cf, - is_inside_staged_cf, - log, - pyir, -) -from .pyir_core import ( # noqa: F401 - MutableValue, - _WatchedM, - _attach_mutable_ref, - _auto_promote_primitive, - _can_create_ref, - _clear_slot_mv, - _create_ref, - _emit_constant_at_ip, - _fresh_wrapper, - _get_instance_attrs, - _get_slot_mv, - _is_boolean_like, - _is_memref_like, - _is_vector_like, - _iter_slot_mvs_for_pyir_read, - _load_as_dsl, - _make_slot_key, - _mlir_type_or_none, - _pyir_assign_simple, - _set_slot_mv, - _slot_storage_available, - _staged_type_changed, -) -from .pyir_corewalk import ( # noqa: F401 - _has_any_staged_content, - _has_decomposable_staged_fields, -) -from .pyir_threading import ( # noqa: F401 - _meta_promote_slot, -) -# -- END explicit imports for the type checker +import array as _array_module +import collections as _collections_module +import dis as _dis_module +import inspect as _inspect_module -import sys - -from typing import TypeVar - -_CM = TypeVar("_CM") +from . import pyir_class_facts as _pyir_class_facts +from .pyir_state import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_core import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_corewalk import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_loop_carry import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_state import _Sentinel def _decompose_tuple( field_key: str, old_tuple: tuple, new_tuple: tuple, - filename: str, - lineno: int, + filename: "str | None", + lineno: "int | None", _visited: set[int], *, owner: Any = None, @@ -103,8 +40,8 @@ def _decompose_tuple( Returns a new tuple with ``pyir_assign``-updated elements. - *owner* / *slot_name* thread the parent's slot key down so per-element - refs key on ``(owner, f"{slot_name}[{i}]")`` -- matching the key + *owner* / *slot_name* carry the parent's slot key down so per-element + refs key on ``(owner, _place_seg_child(slot_name, i))`` -- matching the key convention ``pyir_read``'s tuple recursion uses on the read side. """ if len(old_tuple) != len(new_tuple): @@ -121,9 +58,9 @@ def _decompose_tuple( result: list = [] for i, (old_elem, new_elem) in enumerate(zip(old_tuple, new_tuple)): elem_key = f"{field_key}[{i}]" - elem_slot = f"{slot_name}[{i}]" if has_slot_ctx else None + elem_slot = _place_seg_child(slot_name, i) if has_slot_ctx else None - if _is_staged_value(old_elem) and _can_create_ref(old_elem): + if _is_staged_value(old_elem) and _can_carry_leaf_ref(old_elem): result.append( pyir_assign( elem_key, @@ -151,7 +88,7 @@ def _decompose_tuple( ) elif ( - hasattr(old_elem, "__dict__") + _has_instance_storage(old_elem) and not _is_staged_value(old_elem) and type(old_elem) is type(new_elem) ): @@ -165,6 +102,25 @@ def _decompose_tuple( ) result.append(old_elem) + elif ( + isinstance(old_elem, (int, float, bool, _WatchedM)) + and _is_staged_value(new_elem) + and _can_carry_leaf_ref(new_elem) + ): + # M->S leaf: a primitive loop-init element the body made staged. Route through + # ``pyir_assign`` for a stable loop-carried ref instead of a poison-init ref. + result.append( + pyir_assign( + elem_key, + old_elem, + new_elem, + filename, + lineno, + owner=owner if has_slot_ctx else None, + slot_name=elem_slot, + ) + ) + elif old_elem is None or isinstance(old_elem, (int, float, bool, str, bytes)): result.append(new_elem) @@ -176,15 +132,98 @@ def _decompose_tuple( ) result.append(new_elem) - return tuple(result) + # Preserve a namedtuple's subtype: ``tuple(result)`` would strip field names, + # breaking later ``obj.field`` reads. Use the new tuple as the template. + return _rebuild_tuple_like(new_tuple, result) + + +def _unalias_tuple_leaves(value: tuple) -> tuple: + """LAW-2 at a bare-local tuple binding: give the tuple its OWN ref-supported + leaf wrappers. + + A tuple literal hands its elements straight over from whatever produced + them, so a leaf is routinely the live binding of ANOTHER place -- ``coord = + (i, base, head_idx, zero)`` puts the caller's ``head_idx`` object into the + tuple. The leaf's read-minted row then adopts that very object as its + template, and every later use of ``head_idx`` follows the row (V-2 serves + the template by identity) into a cell stored only where the container was + bound, so a use on a path the container binding does not reach reads the + row's placeholder init. + + A fresh wrapper shares the leaf's SSA, so the tuple's own value is + unchanged; only the object identity the row adopts differs. Leaves whose + type is not ``pyir.ref``-supported (tensors, tensor maps) keep their object: + their carry story is the value-tree walk's, not a leaf cell's, and a tuple + with nothing to unalias is returned as-is. A tuple that DOES own new leaf + wrappers is a new object, so an alias binding of one (``b = a``) no longer + answers ``b is a`` -- the same wrapper-identity break the scalar rule at the + first-def choke already takes, and the reason ``is`` on a tracked value is + refused outright. + """ + leaves = [ + ( + _unalias_tuple_leaves(elem) + if isinstance(elem, tuple) + else ( + _fresh_wrapper(elem) + if _is_staged_value(elem) and _can_create_ref(elem) + else elem + ) + ) + for elem in value + ] + if all(new is old for new, old in zip(leaves, value)): + return value + return _rebuild_tuple_like(value, leaves) + + +def _pyir_decompose_container_entries( + field_key: str, + watched: Any, + filename: "str | None", + lineno: "int | None", + entries: Any, + setter: Any, +) -> None: + """Per-entry M->M decomposition shared by the dict/list arms of + :func:`_decompose_m2m_assign`: identity-equal entries are skipped, a + staged leaf entry routes through the owner-keyed pre-read (publishing the + entry cell under the subscript slot so the assign resolves the SAME cell) + then ``pyir_assign``, and anything else is written back raw. *entries* + yields lazily so each entry is read only after the previous one's write.""" + for _k, _old_e, _new_e in entries: + if _old_e is _new_e: + continue + if _is_staged_value(_old_e) and _can_carry_leaf_ref(_old_e): + _old_e = pyir_read( + f"{field_key}[{_k!r}]", + _old_e, + owner=watched, + slot_name=_k, + ) + setter( + watched, + _k, + pyir_assign( + f"{field_key}[{_k!r}]", + _old_e, + _new_e, + filename, + lineno, + owner=watched, + slot_name=_k, + ), + ) + else: + setter(watched, _k, _new_e) def _decompose_m2m_assign( target_name: str, old_obj: object, new_obj: object, - filename: str, - lineno: int, + filename: "str | None", + lineno: "int | None", *, _visited: set[int] | None = None, ) -> None: @@ -238,11 +277,9 @@ def _decompose_m2m_assign( field_key = f"{target_name}.{attr_name}" - if _is_staged_value(old_field) and _can_create_ref(old_field): - # Thread slot context so the per-field ref identity is keyed - # on ``(old_obj, attr_name)`` -- closes the last loophole - # where compound replacement re-uses the same pyir.ref for - # multiple fields that shared a Python value object. + if _is_staged_value(old_field) and _can_carry_leaf_ref(old_field): + # Carry slot context so per-field ref identity keys on ``(old_obj, attr_name)``, + # avoiding a shared ref for fields that happen to share a Python value. result = pyir_assign( field_key, old_field, @@ -252,7 +289,7 @@ def _decompose_m2m_assign( owner=old_obj, slot_name=attr_name, ) - object.__setattr__(old_obj, attr_name, result) + _pyir_setattr_raw(old_obj, attr_name, result) elif isinstance(old_field, tuple) and isinstance(new_field, tuple): new_tuple = _decompose_tuple( @@ -265,14 +302,65 @@ def _decompose_m2m_assign( owner=old_obj, slot_name=attr_name, ) - object.__setattr__(old_obj, attr_name, new_tuple) + _pyir_setattr_raw(old_obj, attr_name, new_tuple) + + elif ( + isinstance(old_field, dict) + and isinstance(new_field, dict) + and set(dict.keys(old_field)) == set(dict.keys(new_field)) + ): + # Dict field with an IDENTICAL key set: decompose per entry onto the + # dict's subscript legs; the old dict stays the binding carrier. + _watched = _pyir_adopt_dict_value( + old_obj, attr_name, old_field, label=field_key + ) + _pyir_decompose_container_entries( + field_key, + _watched, + filename, + lineno, + entries=( + ( + _k, + dict.__getitem__(_watched, _k), + dict.__getitem__(new_field, _k), + ) + for _k in list(dict.keys(_watched)) + ), + setter=dict.__setitem__, + ) + + elif ( + isinstance(old_field, list) + and isinstance(new_field, list) + and len(old_field) == len(new_field) + ): + # List field with an IDENTICAL length: decompose per element onto the + # list's integer legs; the old list stays the binding carrier. + _watched_l = _pyir_adopt_list_value( + old_obj, attr_name, old_field, label=field_key + ) + _pyir_decompose_container_entries( + field_key, + _watched_l, + filename, + lineno, + entries=( + (_i, list.__getitem__(_watched_l, _i), new_field[_i]) + for _i in range(list.__len__(_watched_l)) + ), + setter=list.__setitem__, + ) elif ( - hasattr(old_field, "__dict__") + _has_instance_storage(old_field) and not _is_staged_value(old_field) and not isinstance(old_field, (int, float, bool, str, bytes, type)) and type(old_field) is type(new_field) + and _has_decomposable_staged_fields(old_field) ): + # Decomposable nested compound: recurse so its staged scalars carry through + # ``pyir.ref`` slots. A non-decomposable field falls through to copy-as-is below. _decompose_m2m_assign( field_key, old_field, @@ -283,18 +371,11 @@ def _decompose_m2m_assign( ) elif old_field is None or isinstance(old_field, (int, float, bool, str, bytes)): - # Meta-primitive field. When this M->M decomposition is - # happening inside staged control flow, the field's value - # MUST stay constant across iterations -- the compiler only - # traces the loop body once, so a "1 -> 2" change is silently - # erased. Catch it here and route to DSLUserCodeError so - # users get a comprehensible diagnostic instead of subtly - # wrong code. - # - # Meta-int wrappers (e.g. ``_WatchedInt``) inherit from ``int`` - # but identify as a distinct ``type()``, so strict type identity - # is unreliable here. ``isinstance(new_field, type(old_field))`` - # or vice versa keeps the check honest across subclass relations. + # A meta-primitive field must stay constant inside staged CF (the + # body is traced once) -- route a change to a diagnostic instead. + + # Meta-int wrappers inherit from ``int`` but report a distinct + # ``type()``: use ``isinstance`` both ways, not type identity. same_numeric_kind = isinstance(new_field, type(old_field)) or isinstance( old_field, type(new_field) ) @@ -308,20 +389,15 @@ def _decompose_m2m_assign( and old_field != new_field ): raise DSLUserCodeError( - f"Meta-primitive field `{type(old_obj).__name__}.{attr_name}` " - f"cannot change across iterations of staged control flow " - f"({old_field!r} -> {new_field!r}). The compiler traces " - f"the body only once and would silently discard later " - f"changes.", + DiagId.PHASE_META_FIELD_CHANGED_IN_CF, filename=filename, lineno=lineno, - suggestion=( - f"Use a DSL Numeric type (e.g. cute.Int32) for " - f"`{attr_name}` so it is tracked by `pyir.ref`, or " - f"hoist the assignment outside staged control flow." - ), + owner_class=type(old_obj).__name__, + attr=attr_name, + old_value=old_field, + new_value=new_field, ) - object.__setattr__(old_obj, attr_name, new_field) + _pyir_setattr_raw(old_obj, attr_name, new_field) else: log().warning( @@ -331,7 +407,7 @@ def _decompose_m2m_assign( attr_name, type(old_field).__name__, ) - object.__setattr__(old_obj, attr_name, new_field) + _pyir_setattr_raw(old_obj, attr_name, new_field) def _pyir_check_no_complex_m2m_call( @@ -340,16 +416,17 @@ def _pyir_check_no_complex_m2m_call( container_repr: str, filename: str, lineno: int, + mutating_values: "list | None" = None, ) -> None: - """Reject in-place mutation of a meta ``list`` / ``dict`` / ``set`` - while inside staged control flow. + """Reject in-place mutation of a meta ``list`` / ``dict`` / ``set`` / + ``collections.deque`` while inside staged control flow. The body of a staged loop / ``scf.if`` is traced exactly once, so a ``a.append(x)`` (or any other mutating method) would silently bake only the first iteration's mutation into the IR -- the per-iteration side effects disappear. Raise a ``DSLUserCodeError`` with a fix-it pointing at slot-backed containers or DSL collection types. A plain - Python ``list`` / ``dict`` / ``set`` is never threaded by the slot + Python ``list`` / ``dict`` / ``set`` is never carried by the slot registry regardless of element type, so the guard fires even when the container already holds staged values (the mutating method itself is not routed through ``pyir_assign``). @@ -375,100 +452,729 @@ def _pyir_check_no_complex_m2m_call( if method_name not in _PYIR_SET_MUTATORS: return kind = "set" + # Only ``add`` with only META values is exempt: the membership-gated + # once-per-key idiom (``if k not in seen: seen.add(k)``) is compile-time + # bookkeeping whose trace-once effect matches Python for every executed + # iteration. Every other set mutator refuses even with meta-only + # values: its trace-once effect is observable post-loop as a phantom + # (a zero-trip loop still sees ``discard``/``clear``/``pop`` applied), + # and ``clear``/``pop`` would slip through vacuously with no values. + if ( + method_name == "add" + and mutating_values is not None + and not any(_is_staged_value(v) for v in mutating_values) + ): + return + elif isinstance(container, _collections_module.deque): + if method_name not in _PYIR_DEQUE_MUTATORS: + return + kind = "deque" else: return - if kind == "list": - # A list mutator (append/extend/insert/pop/...) crossing a region - # boundary changes the list SHAPE; the single trace bakes only pass 0. - raise DSLUserCodeError( - DiagId.CONTAINER_LIST_SHAPE_MUTATED, - var=container_repr, - filename=filename, - lineno=lineno, - ) raise DSLUserCodeError( - f"In-place mutation `{container_repr}.{method_name}(...)` of a " - f"meta Python {kind} is not allowed inside staged control flow. " - f"The compiler traces the body once and would silently discard " - f"the per-iteration mutation.", + DiagId.UNSUP_META_CONTAINER_MUTATION, filename=filename, lineno=lineno, - suggestion=( - f"Build the {kind} before the staged region, or use a DSL " - f"collection type that the compiler can track (e.g. " - f"``cute.struct``-backed buffer for accumulation)." - ), + container=container_repr, + method=method_name, + kind=kind, ) -def pyir_tag_pending_writes(*specs: "tuple[str, Any, Any]") -> None: - """Tag slots in a staged region's SYNTACTIC write-set as pending-store. - - Emitted at the top of ``while_before_block`` for every attr/subscript - the while BODY writes whose owner is readable in the condition scope - (plus the plain-name write_args). The condition evaluates BEFORE the - body's first store lands, so consumers that must decide fold-vs-stage - at that point (``_pyir_while_cond``) need the write-set up front -- - otherwise ``while c.v < 17: c.v += 1`` bakes ``scf.condition(true)`` - from the trace-time fold: an unkillable runtime hang. Each spec is - ``(dotted_name, owner_or_None, slot_or_None)``. - """ - if not is_inside_staged_cf() or not is_auto_m2s_enabled(): +def _pyir_freeze_staged_container_inserts( + container: object, + method_name: str, +) -> None: + """Freeze foreign-slot-bound STAGED elements right after an insert into a meta + container in a constexpr loop, pinning each at its insertion-time SSA.""" + if pyir is None or not is_inside_constexpr_loop(): return - fn_id = _current_fn_id() - for dotted, owner, slot_name in specs: - slot = _make_slot_key(dotted, owner, slot_name, fn_id) - if slot is not None: - _slot_pending_store.add(slot) - - -def pyir_promote_loop_body_arg(target_name: str, current_value: object) -> object: - """Auto-promote a Python-primitive write_arg at loop-body entry. + if isinstance(container, list): + if method_name not in _PYIR_LIST_INSERT_MUTATORS: + return + for idx, elem in enumerate(container): + if _carries_foreign_slot_binding(elem): + container[idx] = _freeze_foreign_slot_binding( + elem, "constexpr-loop container insert" + ) + elif isinstance(container, dict): + if method_name not in _PYIR_DICT_INSERT_MUTATORS: + return + for key, elem in list(container.items()): + if _carries_foreign_slot_binding(elem): + container[key] = _freeze_foreign_slot_binding( + elem, "constexpr-loop container insert" + ) - If a loop-body reads a Python primitive (bool/int/float) that is later - re-stored in the same iteration, the read site bakes the trace-time - constant instead of loading from the loop-carried ref. Calling - ``pyir_read`` at body entry forces the slot to materialise so - subsequent reads inside the body load from it. - A genuinely staged SCALAR write_arg (created outside the loop and read - before its first in-body write) is reloaded from its ref here so the body - observes the loop-carried value -- gated on the value being staged (WS4-B - below), NOT on AUTO_M2S. Memref-backed and vector values are left as-is. - """ +def pyir_promote_loop_body_arg( + target_name: str, + current_value: object, + *, + bare_first_use: bool = False, + carried_attr_paths: "tuple[tuple[str, ...], ...]" = (), + whole_rebound: bool = False, +) -> object: + """Materialise a loop-carried write_arg's ref at body entry so body reads + load the carry, not the pre-loop SSA; ``bare_first_use`` = plain-use first read.""" if not is_inside_staged_cf(): return current_value - if not is_auto_m2s_enabled() and ( - not _is_staged_value(current_value) or _is_memref_like(current_value) + # ``scf.while`` after-block re-load of the declared carried attr legs: the + # before-block SSA the condition promotion rebound does not dominate here. + if ( + carried_attr_paths + and current_value is not None + and not isinstance(current_value, (bool, int, float, str, bytes)) + and _has_instance_storage(current_value) + ): + _promote_carried_attr_legs( + target_name, current_value, carried_attr_paths, force_promote_meta=False + ) + if isinstance(current_value, tuple): + # Recurse per element under the indexed slot key so each leaf carries as + # its own iter_arg; ``bare_first_use`` gates per leaf, not wholesale. + return _rebuild_tuple_like( + current_value, + [ + pyir_promote_loop_body_arg( + f"{target_name}[{i}]", elem, bare_first_use=bare_first_use + ) + for i, elem in enumerate(current_value) + ], + ) + if ( + whole_rebound + and (type(current_value) is list or isinstance(current_value, _WatchedList)) + and current_staged_cf_depth() <= 1 ): + # Loop-carried WHOLE-rebound list: promote each meta-primitive element + # through its owner-keyed subscript cell (outermost loop only; raw store-back). + for _i in range(list.__len__(current_value)): # type: ignore[arg-type] + _elem = list.__getitem__(current_value, _i) # type: ignore[arg-type] + _init = _elem.python_value if isinstance(_elem, _WatchedM) else _elem + if type(_init) not in (bool, int, float): + continue + _fresh = pyir_read( + f"{target_name}[{_i}]", + _elem, + owner=current_value, + slot_name=_i, + force_promote_meta=True, + ) + if _fresh is not _elem: + list.__setitem__(current_value, _i, _fresh) # type: ignore[arg-type] + return current_value + if ir is not None and isinstance(current_value, ir.Value): + _pyir_refuse_stale_raw_carry(target_name, current_value, "loop-carried") + if isinstance(current_value.type, (ir.IntegerType, ir.FloatType)): + # A raw scalar SSA write_arg is a loop-carried place: mint its + # cell (the body's write choke stores raw SSA rebinds through the + # same row) and serve body reads from it, so consumers read the + # carried arg, not the pre-loop trace-time SSA. + return _pyir_promote_raw_scalar_carry(target_name, current_value) return current_value if _is_staged_value(current_value): - # Already-staged scalar: reload from its ref at body entry so in-body reads - # emitted BEFORE the first write observe the loop-carried value instead of - # the cached trace-entry SSA. Restricted to ref-capable, non-vector, - # non-memref scalars -- the ref lands at the value's dominating def, so no - # new placement order is introduced. - if ( - _can_create_ref(current_value) - and not _is_vector_like(current_value) - and not _is_memref_like(current_value) + if not _can_create_ref(current_value): + return current_value + # Boolean/Vector args route through the generic read unconditionally: + # it creates-or-reuses the place cell and returns a dominating load. + if _is_boolean_like(current_value) or _is_vector_like(current_value): + return pyir_read(target_name, current_value) + # A staged scalar with NO accessible ref and a plain-use first read + # must materialise the ref so the read loads the carried value. + if bare_first_use and not _has_accessible_loop_carried_ref( + target_name, current_value ): return pyir_read(target_name, current_value) return current_value + # Value-tree compounds are already carried per leaf; pass through. if not isinstance(current_value, (bool, int, float)): return current_value - return pyir_read(target_name, current_value) + return _promote_loop_carried_meta(target_name, current_value) + + +def _pyir_record_fresh_entry_birth(owner: Any, slot_name: Any, entry: Any) -> None: + """Record the current block as the F-BIRTHPOS fact of a fresh container + entry's meta leaves (tuples recurse per leaf). Recording only, no ref: + a never-promoted entry stays a plain meta value.""" + if isinstance(entry, tuple): + for i, elem in enumerate(entry): + _pyir_record_fresh_entry_birth(owner, _place_seg_child(slot_name, i), elem) + return + payload = _pyir_unwrap_meta_primitive(entry) + if payload is None or type(payload) not in (bool, int, float): + return + try: + slot = _make_slot_key(None, owner, slot_name) + block = ir.InsertionPoint.current.block + except Exception: + return + if slot is None or slot in _slot_refs or slot in _slot_first_def_block: + return + _slot_first_def_inside_cf[slot] = True + _slot_first_def_depth[slot] = current_staged_cf_depth() + _slot_first_def_block[slot] = block + _pyir_keepalive_generation_obj(owner) + + +def _pyir_fresh_container_init( + label: str, + container: Any, + filename: "str | None", + lineno: "int | None", + fresh_paths: "frozenset[tuple]", + _prefix: tuple = (), +) -> None: + """Per-iteration re-init of a fresh in-region container's entry cells (the + scalar in-region first-def rule); recursion follows declared fresh paths only.""" + items = ( + [(k, dict.__getitem__(container, k)) for k in dict.keys(container)] + if isinstance(container, _WatchedDict) + else list(enumerate(list.__iter__(container))) + ) + for k, e in items: + entry_label = f"{label}[{k!r}]" + if isinstance(e, (_WatchedDict, _WatchedList)): + # Nested containers were adopted eagerly at the root adoption; + # descend only along a DECLARED nested-construction path. + child_path = _prefix + (k,) + if child_path in fresh_paths: + _pyir_fresh_container_init( + entry_label, e, filename, lineno, fresh_paths, child_path + ) + continue + if not ( + _is_staged_value(e) + and _can_carry_leaf_ref(e) + and not _is_compound_single_leaf(e) + ): + # Meta entry (scalar or tuple of scalars): record the construction + # site as the entry place's birth, so a later promotion seeds a + # real store here (per-iteration re-init) instead of hoisting the + # init to function entry, which would carry the value across + # iterations. + _pyir_record_fresh_entry_birth(container, k, e) + continue + log().info( + "[pyir_assign] '%s' fresh-container binding inside staged CF → " + "per-entry in-region init", + entry_label, + ) + init_e = pyir_assign( + entry_label, + None, + e, + filename, + lineno, + owner=container, + slot_name=k, + ) + if isinstance(container, _WatchedDict): + dict.__setitem__(container, k, init_e) + else: + list.__setitem__(container, k, init_e) + + +def _pyir_route_binding_to_live_place_cells( + owner: Any, + slot_name: Any, + new_value: Any, + _visited: "set[int] | None" = None, +) -> None: + """Complete a binding made outside staged CF: store each staged leaf into + its place's live cell and re-alias, so reads and writes keep ONE cell.""" + if _visited is None: + _visited = set() + # Top-level staged scalar bound at an explicit slot. + if owner is not None and slot_name is not None and _is_staged_value(new_value): + _pyir_adopt_live_place_cell(owner, slot_name, new_value) + return + # Compound: complete each staged leaf against ITS place cell. + if ( + new_value is None + or _is_staged_value(new_value) + or isinstance(new_value, (int, float, bool, str, bytes, type)) + or not _has_instance_storage(new_value) + ): + return + if id(new_value) in _visited: + return + _visited.add(id(new_value)) + for attr_name in _get_instance_attrs(new_value): + try: + field = getattr(new_value, attr_name) + except AttributeError: + continue + if _is_staged_value(field) and _can_carry_leaf_ref(field): + _pyir_adopt_live_place_cell(new_value, attr_name, field) + elif isinstance(field, tuple): + for i, elem in enumerate(field): + if _is_staged_value(elem) and _can_carry_leaf_ref(elem): + _pyir_adopt_live_place_cell( + new_value, _place_seg_child(attr_name, i), elem + ) + elif ( + _has_instance_storage(field) + and not _is_staged_value(field) + and not isinstance(field, (int, float, bool, str, bytes, type)) + and _has_decomposable_staged_fields(field) + ): + _pyir_route_binding_to_live_place_cells( + None, None, field, _visited=_visited + ) + + +# Dynamic extent of the machinery's whole-object M->M decomposition: its +# per-field replays re-state a binding the choke already judged, so the +# declared-surface wall must not re-judge them as direct field writes. +_PYIR_M2M_DECOMPOSE_DEPTH: "list[int]" = [0] + + +def _pyir_surface_wall_governed_value(value: Any) -> bool: + """Whether *value* has a choke-governed carry story as an attr-write + payload: meta primitives (trace-time accounting), staged / raw-``ir.Value`` + leaves (per-field ref carry), containers and tuples (adoption / + decomposition), and the watched-meta wrapper. Everything else is an + opaque compound whose predicated rebind nothing carries.""" + return ( + value is None + or isinstance( + value, + (bool, int, float, str, bytes, type, tuple, list, dict, set, frozenset), + ) + or isinstance(value, ir.Value) + or isinstance(value, _WatchedM) + or _is_staged_value(value) + ) + + +def _pyir_judge_declared_surface_write( + owner: Any, + slot_name: Any, + old_value: Any, + new_value: Any, + filename: "str | None", + lineno: "int | None", +) -> None: + """Declaration wall for value-tree-protocol owners: replacing an OPAQUE + COMPOUND held by a field the owner's OWN ``__extract_mlir_values__`` + never touches, inside dynamic staged CF, refuses loudly. Such a field + was declared constant context and its payload has no carry story, so + the trace-time replacement could not follow the region's runtime + predicate (it would bake unconditionally). + + The declared write model stays admitted: ctor-birth first-defs + (``old_value is None``), meta-primitive accounting, staged / raw-leaf + per-field carry, container adoption, and the machinery's own + per-field decompose replays. + + Companion arm: a wrapper subclass that overrides ``__hash__`` beneath + the DSL value type answers hashing without the wrapper's consumption + witness -- the same declaration break, judged at the same choke. + + An underivable surface produces no fact and admits the write; a + constexpr scope resolves at trace time and is exempt like every + meta-mutation judgment.""" + # A ``_PlaceSeg`` slot names a container ELEMENT under a field; only the + # plain-str FIELD binding is a declared-surface judgment. + if owner is None or not isinstance(slot_name, str): + return + if _PYIR_M2M_DECOMPOSE_DEPTH[0]: + return # machinery per-field replay of a judged whole-object rebind + if not is_inside_staged_cf() or is_inside_constexpr_loop(): + return + if not _implements_dynamic_expression(owner): + return + cls = type(owner) + definer = _pyir_class_facts.staged_wrapper_hash_override(cls) + if definer is not None: + from .pyir_call_boundary import _pyir_boundary_module_is_user + + if _pyir_boundary_module_is_user(getattr(definer, "__module__", None)): + raise DSLUserCodeError( + DiagId.OWNER_DECLARED_SURFACE_VIOLATION, + filename=filename, + lineno=lineno, + obj=cls.__name__, + violation=( + f"`{definer.__qualname__}` overrides `__hash__` beneath " + "the DSL value type, hiding its identity witness" + ), + ) + # Ctor-birth first-def, or a payload kind with a choke-governed carry + # story: the declared write model owns it. + if old_value is None: + return + if _pyir_surface_wall_governed_value( + old_value + ) or _pyir_surface_wall_governed_value(new_value): + return + surface = _pyir_class_facts.declared_extract_surface(cls) + if surface is not None and slot_name not in surface: + raise DSLUserCodeError( + DiagId.OWNER_DECLARED_SURFACE_VIOLATION, + filename=filename, + lineno=lineno, + obj=cls.__name__, + violation=( + f"`{slot_name}` is not part of that declaration and is written here" + ), + ) + + +def _pyir_carry_meta_compound_region_rebind( + old_value: Any, + new_value: Any, + binding_slot: Any, + cur_cf_depth: int, + target_name: Any, + filename: "str | None", + lineno: "int | None", +) -> Any: + """Carry a REGION-CROSSING whole-object ctor rebind of an all-meta- + primitive-field compound by promoting its numeric fields per-field. + + A compound with no staged content passes through ``pyir_assign`` as a + harmless trace-time replacement -- but when the binding was born at a + SHALLOWER staged depth (it carries across iterations of the enclosing + for/while/if) and the new object's fields are all meta primitives, the + trace-once body would bake the fields at their first-iteration values. + Instead of refusing, promote each bool/int/float field to a staged slot + keyed to the OLD object -- the SAME place every earlier ``obj.attr`` + read recorded its bakes under, so ``_meta_promote_slot`` rewrites those + bakes to loads -- store the new field's SSA through it, and keep the + OLD object as the binding carrier (m2m polarity): the compound twin of + bare-int D1 promotion, threading each field as a carried phi the way + the L191 staged-field whole-object rebind does. + + Returns the carrier (the OLD object) when the rebind was carried, or + ``None`` when the passthrough stays legal: identity rebinds, constexpr + unrolls (realized at trace time), in-region-born bindings (trace-local + rebind), unknown birth depth (uninstrumented first-def, e.g. a callee + parameter), objects with any non-primitive field (opaque/staged carry + stories own those), and same-class all-fields-equal rebinds (a fixed + point of meta state -- re-running the body cannot produce different + fields). + + Raises ``PHASE_META_FIELD_CHANGED_IN_CF`` only for a CHANGED field with + no sound carry story: a field the old object does not carry as a meta + primitive (missing / None / nested object -- "unchanged" is unprovable + there), str/bytes payloads, cross-type payload drift, + a new value with no eager-IR derivation (a plain payload or a + payload-domain fold would store a trace-time CONSTANT), a field already + consumed as trace-time structure (folded if/while tests, coercions, + comparisons -- the bake is not retargetable), a cross-class rebind, or + a promotion the meta-value table does not admit. + """ + if old_value is new_value or is_inside_constexpr_loop(): + return None + first_depth = ( + _slot_first_def_depth_any.get(binding_slot) + if binding_slot is not None + else None + ) + if first_depth is None or first_depth >= cur_cf_depth: + return None + attrs = _get_instance_attrs(new_value) + if not attrs or not all( + isinstance(getattr(new_value, a), (bool, int, float, str, bytes)) for a in attrs + ): + return None + + def _refuse(attr: str, old_field: Any, new_field: Any) -> None: + raise DSLUserCodeError( + DiagId.PHASE_META_FIELD_CHANGED_IN_CF, + filename=filename, + lineno=lineno, + owner_class=type(old_value).__name__, + attr=attr, + old_value=old_field, + new_value=new_field, + ) + + changed = [] + _absent = object() + for attr in attrs: + new_field = getattr(new_value, attr) + old_field = getattr(old_value, attr, _absent) + # Compare RAW payloads: either side may arrive as a ``_WatchedM``, + # and a wrapper comparison would coerce through the watch dunders. + # Primitives compare safely across types (no user dunders), so a + # type drift with a value change (int 0 -> float 0.5) counts as a + # change; equal-by-value pairs (True -> 1) stay a fixed point. + old_py = getattr(old_field, "_pyir_raw_payload", old_field) + new_py = getattr(new_field, "_pyir_raw_payload", new_field) + if not isinstance(old_py, (bool, int, float, str, bytes)): + # The field APPEARED, or its OLD kind is not a meta primitive + # (None / nested object): "unchanged" is unprovable and there + # is no numeric cell to thread through -- fail closed instead + # of letting the rebind pass as a fixed point and bake. + _refuse(attr, "" if old_field is _absent else old_py, new_py) + if old_py != new_py: + # A textual payload has no numeric cell to thread through. + if isinstance(old_py, (str, bytes)) or isinstance(new_py, (str, bytes)): + _refuse(attr, old_py, new_py) + # A cross-type numeric drift (int 0 -> float 0.5) has no single + # SSA element type for the carried cell -- fail closed. + if type(old_py) is not type(new_py): + _refuse(attr, old_py, new_py) + # The carried store must RE-DERIVE from the promoted loads on + # every iteration: only an eager-IR derivation (``_binop_ir`` + # chains) is rewrite-reachable. A plain payload or a + # payload-domain wrapper (comparison results, ``not_`` folds) + # would store the trace-time CONSTANT of a folded decision -- + # fail closed. + if ( + not isinstance(new_field, _WatchedM) + or getattr(new_field, "_cached_ir", None) is None + ): + _refuse(attr, old_py, new_py) + changed.append((attr, old_py, new_py)) + if type(old_value) is not type(new_value): + # Cross-class all-meta rebind: the field sets need not align and the + # bound methods differ across iterations, so no per-field carry story + # exists. Equal fields are a fixed point of the FIELDS, not of the + # object (the traced body baked the old class's methods) -- refuse + # before the passthrough, on a witness field carrying the class names + # when no field changed. + attr, old_py, new_py = ( + changed[0] + if changed + else ( + next(iter(attrs)), + type(old_value).__name__, + type(new_value).__name__, + ) + ) + _refuse(attr, old_py, new_py) + if not changed: + return None # same-class all-fields-equal fixed point: passthrough stays legal + changed_names = {attr for attr, _, _ in changed} + for attr, old_py, new_py in changed: + # A changed field already consumed as trace-time STRUCTURE (an + # `if obj.attr` predicate fold, a user-frame coercion or comparison, + # `__index__`, ...) has no retargetable SSA bake: the carry would + # keep the stale decision on every iteration -- refuse on the + # witness the bare-local D1 arm refuses on. + if ( + _PYIR_STRUCTURAL_META_CONSUMPTIONS.get( + _make_slot_key(None, old_value, attr) + ) + is not None + ): + _refuse(attr, old_py, new_py) + + # m2m polarity: the binding keeps the old object as carrier; the + # generation bump lets a pre-rebind alias capture detect staleness. + _gen_rebind_rec = _pyir_record_generation_rebind( + binding_slot, + old_value, + new_value, + target_name, + filename, + lineno, + allow_empty_cells=True, + ) + for attr in attrs: + new_field = getattr(new_value, attr) + old_field = getattr(old_value, attr, None) + if not isinstance(old_field, (bool, int, float)) or isinstance( + new_field, (str, bytes) + ): + # Textual field: a changed textual field refused above (and a + # missing/non-primitive old field refused in the change scan), + # so this is a fixed point -- copy the payload. + _pyir_setattr_raw(old_value, attr, new_field) + continue + field_name = f"{target_name}.{attr}" + # Owner-keyed pre-read FIRST: it wraps the payload as the watched + # meta value of the SAME place the body's earlier reads baked + # against (else pyir_assign would mint a TWIN cell). + _old_f = pyir_read(field_name, old_field, owner=old_value, slot_name=attr) + carried = pyir_assign( + field_name, + _old_f, + new_field, + filename, + lineno, + owner=old_value, + slot_name=attr, + ) + field_slot = _make_slot_key(field_name, old_value, attr) + if field_slot is None or field_slot not in _slot_refs: + # Promotion not admitted: fail closed on a changed field; an + # unchanged field stays a legal meta fixed point. + if attr in changed_names: + _refuse(attr, old_field, new_field) + _pyir_setattr_raw(old_value, attr, new_field) + continue + _pyir_setattr_raw(old_value, attr, carried) + # Raw-ref born-at-rebind cell: a pre-rebind alias capture reading + # this field resolves the ONE live row and must refuse, exactly + # like the staged-field m2m rebind's superseded-generation reads. + _pyir_record_rebind_ref_cell(_gen_rebind_rec, attr, _slot_refs.get(field_slot)) + _pyir_complete_generation_rebind_cells(_gen_rebind_rec, old_value) + log().info( + "[pyir_assign] '%s' region-crossing all-meta compound rebind -> " + "per-field promotion (%d changed field(s)); old object carries", + target_name, + len(changed_names), + ) + return old_value + + +def _unwrap_meta_or_self(value: Any) -> Any: + """The raw meta payload of *value*, or *value* itself when not a + watched/wrapped meta primitive.""" + payload = _pyir_unwrap_meta_primitive(value) + return value if payload is None else payload + + +_GATE_META_SCALARS = (bool, int, float, str, bytes, type(None)) + + +def _all_meta_content(value: Any, seen: "set[int]") -> bool: + """True when *value* is a meta scalar or a tuple/list/dict/set/plain + object holding only such content, recursively. Staged leaves and + value-protocol objects are not all-meta: a carry path owns them.""" + value = _unwrap_meta_or_self(value) + if isinstance(value, _GATE_META_SCALARS): + return True + if _is_staged_value(value) or _implements_dynamic_expression(value): + return False + if id(value) in seen: + # A shared sub-object or a cycle adds no new content; the first + # visit judges it (a changed cycle fails closed in the compare). + return True + seen.add(id(value)) + if isinstance(value, (tuple, set, frozenset)): + return all(_all_meta_content(e, seen) for e in tuple(value)) + if isinstance(value, list): + return all(_all_meta_content(e, seen) for e in list.__iter__(value)) + if isinstance(value, dict): + return all(_all_meta_content(v, seen) for v in dict.values(value)) + names = _get_instance_attrs(value) + if not names: + return False + return all(_all_meta_content(getattr(value, n), seen) for n in names) + + +def _meta_content_tree(value: Any, seen: "set[int]") -> Any: + """*value* as a plain comparable tree: leaves unwrapped, containers + as builtins, an object as its class plus field tree. Lets the gate + compare content by value where objects compare by identity.""" + value = _unwrap_meta_or_self(value) + if isinstance(value, _GATE_META_SCALARS): + return value + if id(value) in seen: + return "" + seen.add(id(value)) + if isinstance(value, tuple): + return tuple(_meta_content_tree(e, seen) for e in value) + if isinstance(value, list): + return [_meta_content_tree(e, seen) for e in list.__iter__(value)] + if isinstance(value, (set, frozenset)): + return set(value) + if isinstance(value, dict): + return {k: _meta_content_tree(v, seen) for k, v in dict.items(value)} + return ( + type(value), + { + n: _meta_content_tree(getattr(value, n), seen) + for n in _get_instance_attrs(value) + }, + ) + + +def _refuse_meta_compound_region_rebind( + owner: Any, + slot_name: Any, + old_value: Any, + new_value: Any, + cur_cf_depth: int, + target_name: Any, + filename: "str | None", + lineno: "int | None", +) -> None: + """Refuse a rebind of an attr-held all-meta compound (tuple, list, + dict, set, or plain object, nested to any depth) whose content + changed inside staged CF: no cell carries it, so the trace-once body + would keep the first iteration's content. Unchanged content, region-born bindings, + constexpr unrolls, staged content, and container-element + (``_PlaceSeg``) slots pass through. Value-protocol owners are the + declared-surface wall's business (including its deliberate fail-open + cases), so this gate skips them.""" + if owner is None or not isinstance(slot_name, str) or is_inside_constexpr_loop(): + return + if _implements_dynamic_expression(owner): + return + kind = None + if isinstance(old_value, tuple) and isinstance(new_value, tuple): + kind = "tuple" + elif isinstance(old_value, list) and isinstance(new_value, list): + kind = "list" + elif isinstance(old_value, dict) and isinstance(new_value, dict): + kind = "dict" + elif isinstance(old_value, (set, frozenset)) and isinstance( + new_value, (set, frozenset) + ): + kind = "set" + elif type(old_value) is type(new_value) and _get_instance_attrs(old_value): + # Fields may nest compounds (a tuple field, an object field): + # judge the whole content tree, not just scalar fields. + if _all_meta_content(old_value, set()) and _all_meta_content(new_value, set()): + kind = "object" + if kind is None: + return + if _has_any_staged_content(old_value) or _has_any_staged_content(new_value): + return # staged content has its own carry/refusal paths + slot = _make_slot_key(target_name, owner, slot_name) + first_depth = _slot_first_def_depth_any.get(slot) if slot is not None else None + if first_depth is not None and first_depth >= cur_cf_depth: + return # region-born binding: a trace-local rebind + try: + if kind == "object": + changed = _meta_content_tree(old_value, set()) != _meta_content_tree( + new_value, set() + ) + else: + changed = bool(old_value != new_value) + except Exception: + changed = True # incomparable payloads: fail closed + if not changed: + return # content unchanged: re-tracing reproduces it + if kind == "set": + raise DSLUserCodeError( + DiagId.CONTAINER_SET_REBUILT_IN_CF, + filename=filename, + lineno=lineno, + var=target_name, + old_value=old_value, + new_value=new_value, + ) + raise DSLUserCodeError( + DiagId.CONTAINER_META_REBUILT_IN_CF, + filename=filename, + lineno=lineno, + var=target_name, + kind=kind, + old_value=old_value, + new_value=new_value, + ) def pyir_assign( target_name: Any, old_value: Any, new_value: Any, - filename: str, - lineno: int, + filename: "str | None", + lineno: "int | None", *, owner: Any = None, slot_name: Any = None, + fresh_binding: bool = False, + fresh_paths: tuple = (), + rhs_ctor: Any = None, ) -> Any: """Called by AST-inserted code after every ``=`` and ``+=``. @@ -479,10 +1185,11 @@ def pyir_assign( - Returns ``new_value`` with ``_mutable_ref`` attached (NO load). 3. Outside CF: returns new_value unchanged. - The deferred-load invariant: ``pyir_assign`` never emits - ``pyir.load``. Loads are deferred to ``pyir_read`` or - ``_pyir_auto_load_arg`` at the use site, ensuring the loaded SSA - value is created at the correct insertion point (outside CF). + Value-keyed local writes defer their loads to ``pyir_read`` / + ``_pyir_auto_load_arg`` at the use site (the loaded SSA is created at + the correct insertion point); the D1 owner/slot arms and the + store-through paths DO return fresh ``pyir.load`` products so the + binding the caller re-binds is cell-backed. *owner* / *slot_name* optionally identify the storage slot (e.g. ``(obj, "loop_desc")`` or ``(container, "key")``). When both are @@ -490,9 +1197,16 @@ def pyir_assign( on the value object -- preventing shared-value aliasing bugs where three attributes initialised from the same Python object collapse onto a single ``pyir.ref``. When either is ``None`` the function - falls back to the legacy value-keyed path (preserving local-name + falls back to the value-keyed path (preserving local-name semantics). + *rhs_ctor* is the re-loaded callee of a DIRECT-call RHS (``x = Ctor(...)``, + Name-rooted load chain): with standard allocation semantics it witnesses + that the statement CONSTRUCTED the bound object, so its raw meta-primitive + fields record their in-region first-def facts at the binding (see + :func:`_pyir_record_fresh_object_leaf_first_defs`). ``None`` for any + other RHS shape. + Simplified ``pyir_assign(owner, key, value, filename, lineno)`` entry point: detected by *target_name* not being a string (it's the owner object) and dispatched to the slot-registry-only path; the location @@ -502,27 +1216,123 @@ def pyir_assign( # Simplified (owner, key, value) dispatch -- used by the slot # registry unit test and by external (owner, key) callers. The # location args are still required (see signature) so this branch - # never observes ``None`` for either; we just don't thread them + # never observes ``None`` for either; we just don't carry them # through ``_pyir_assign_simple`` (which is registry-only and emits # no diagnostics). if not isinstance(target_name, str): return _pyir_assign_simple(target_name, old_value, new_value) - # Register a placeholder for the slot in _slot_mvs keyed by - # structural identity. ``_get_slot_mv`` consults the registry for - # registry-backed owners (dict/list, which have no ``__dict__``); - # ``__dict__``-backed owners use the tier-1/tier-2 stores below. - # Authoritative storage stays in _set_slot_mv / _get_slot_mv below. + # F-MEMORY: an element of a memref-backed owner lives in staged memory and + # the owner's accessors emit its stores; memory, not a row, is the authority. + if ( + owner is not None + and slot_name is not None + and not isinstance(slot_name, (str, _PlaceSeg)) + and _is_memref_like(owner) + ): + return new_value + # Declaration wall: a value-tree-protocol owner's observed field write in + # dynamic staged CF must stay on its declared (extraction-touched) surface. + _pyir_judge_declared_surface_write( + owner, slot_name, old_value, new_value, filename, lineno + ) + # F-SPEC write amendment: a trace-observed persistent-state write re-bakes + # its recorded rows so re-entry verifies the trace-exit value. if owner is not None and slot_name is not None: - slot_key = _make_slot_key(None, owner, slot_name) - if slot_key not in _slot_mvs: - placeholder = MutableValue.__new__(MutableValue) - object.__setattr__(placeholder, "_value", new_value) - object.__setattr__(placeholder, "_type", type(new_value)) - object.__setattr__(placeholder, "_ref", None) - object.__setattr__(placeholder, "_ref_context_id", None) - object.__setattr__(placeholder, "_load_version", 0) - _slot_mvs[slot_key] = placeholder + # Host-restore witness: the pre-write scalar binding this place returns + # to at trace close if the write leaves it holding a staged wrapper. + _pyir_record_host_restore(owner, slot_name, old_value) + _pyir_spec_record_write( + _pyir_read_place(target_name, owner, slot_name), new_value + ) + # F-CEPLACE: the assign choke maintains the binding-BIRTH ownership fact + # for bare locals; every later key derivation in this call consumes it. + # A non-None *old_value* witnesses a rebind of an existing binding (the + # instrumented form read the bound value first), which continues the one + # logical binding; only a first-def can birth a new one. + if owner is None and slot_name is None: + _ce_note_local_assign(target_name, continues_binding=old_value is not None) + # A staged write of a place a plain callee already READ through a closure + # cell in staged CF cannot be retargeted; refuse loudly (declared fact). + if _PYIR_BOUNDARY_META_CELL_READS and owner is None and is_inside_staged_cf(): + _cr_key = _make_slot_key(target_name, owner, slot_name) + _cr = ( + _pyir_boundary_cell_read_for_local(_cr_key) if _cr_key is not None else None + ) + if _cr is not None: + _cr_name, _cr_file, _cr_line = _cr + raise DSLUserCodeError( + DiagId.BOUNDARY_CLOSURE_READ_THEN_WRITTEN, + filename=filename, + lineno=lineno, + name=str(target_name), + def_file=_cr_file, + def_line=_cr_line, + ) + # Candidate-holder registry: an instrumented binding sights the slot owner + # and both bound objects -- each is a potential captured-holder walk root. + _pyir_register_candidate_holder(owner) + _pyir_register_candidate_holder(old_value) + _pyir_register_candidate_holder(new_value) + # Binding-choke adoption: a plain dict at a KNOWN holder slot adopts with the + # holder replaced; owner-less bindings adopt only for FRESH constructions. + if ( + type(new_value) is dict + and not isinstance(old_value, dict) + and _WATCHED_DICT_READ_HOOK[0] is not None + ): + if owner is not None and slot_name is not None: + new_value = _pyir_adopt_dict_value( + owner, slot_name, new_value, label=str(target_name) + ) + elif fresh_binding: + new_value = _pyir_adopt_dict_value( + None, None, new_value, label=str(target_name) + ) + # List sibling of the binding-choke adoption (known holder slot or owner-less + # fresh construction; a list-over-list rebind belongs to the decomposition). + if ( + type(new_value) is list + and not isinstance(old_value, list) + and _WATCHED_LIST_READ_HOOK[0] is not None + ): + if owner is not None and slot_name is not None: + new_value = _pyir_adopt_list_value( + owner, slot_name, new_value, label=str(target_name) + ) + elif fresh_binding: + new_value = _pyir_adopt_list_value( + None, None, new_value, label=str(target_name) + ) + # Record stable-anchor rooting for a dotted attr access; ``alias_from=old_value`` + # lets a NEW object at a rooted place adopt the old rooting (anti-twin). + _register_place_prefix( + target_name, owner, slot_name, old_value, new_value, alias_from=old_value + ) + + # A write through a stamped superseded generation would corrupt the one + # live cell -> diagnostic; every compound binding is recorded for rebinds. + if _SUPERSEDED_GENERATIONS and owner is not None: + _pyir_check_superseded_owner_write(owner, slot_name, new_value, target_name) + _pyir_record_compound_binding( + target_name, owner, slot_name, new_value, filename, lineno + ) + # Region-conditional attr first-def maintenance: a reassignment may widen + # or discharge an earlier in-region first-def record (read-before-set). + if old_value is not None and owner is not None and _PYIR_CF_ATTR_FIRST_DEFS: + _pyir_update_cf_attr_first_def_on_write(owner, slot_name) + + # Binding-depth bookkeeping for the within-body-local gate below: ``_prior_binding_depth`` is the depth + # where ``old_value`` was last written; refreshed to THIS write. Tracked only inside staged CF. + _binding_slot = None + _prior_binding_depth = None + _cur_cf_depth = current_staged_cf_depth() + if is_inside_staged_cf(): + _binding_slot = _make_slot_key(target_name, owner, slot_name) + if _binding_slot is not None: + if old_value is not None: + _prior_binding_depth = _slot_binding_depth.get(_binding_slot) + _slot_binding_depth[_binding_slot] = _cur_cf_depth log().info( "[pyir_assign] '%s' old=%s new=%s (%s:%d)", target_name, @@ -532,232 +1342,119 @@ def pyir_assign( lineno, ) - # First-time definition — create ref eagerly if inside staged CF. - # Only create refs for scalar-like types (Numeric, not Boolean). - # Multi-element types (Vector), boolean types (i1), and types with - # complex ir_value() are left to the reassignment path to handle. - # - # Anti-aliasing: ``a = b = c = seed`` binds three Python locals to - # the same value object. ``ast.Name`` targets carry no storage - # owner, so ``pyir_assign``/``pyir_read`` fall back to the - # value-keyed ``_mutable_ref`` cache. Without a fresh wrapper per - # first-def, the three locals would alias onto ``seed._mutable_ref`` - # and collapse onto a single ``pyir.ref``. Returning a fresh - # wrapper per first-def gives each local its own attachment slot. if old_value is None: - # The default first-def gate excludes booleans because a free- - # standing ``cond = i > 2`` is almost always a one-shot if-test - # and the dead ref breaks PYIRToSCF. But when the caller supplied - # slot context (e.g. ``self._is_valid_tile = Boolean(valid)`` in - # ``@cute.jit __init__``), the slot store records the ref so - # downstream ``pyir_read('container', container)`` can refresh it - # across staged CF boundaries. Slot-keyed booleans are therefore - # allowed to take refs. - # - # CRITICAL: keep the gate as a single ``and`` chain so the early - # checks short-circuit BEFORE we call ``_is_boolean_like`` / - # ``_is_vector_like``. Those helpers call ``value.ir_value()`` - # which materialises (and caches) the leaf ``arith.constant`` for - # ``_WatchedInt`` -- if we evaluate them on a meta primitive we - # pin the constant at the wrong insertion point and break the - # sibling-region constant cache. - if ( - is_inside_staged_cf() - and _is_staged_value(new_value) - and _can_create_ref(new_value) - and not _is_vector_like(new_value) - and ( - # Boolean blocked unless slot context supplies an - # explicit storage identity. - not _is_boolean_like(new_value) - or ( - owner is not None - and slot_name is not None - and _slot_storage_available(owner) - ) - ) - ): - log().info( - "[pyir_assign] '%s' first def inside staged CF → create ref", - target_name, - ) - new_value = _fresh_wrapper(new_value) - fresh_mv = _create_ref(new_value) - fresh_mv.store(new_value) - _attach_mutable_ref( - new_value, fresh_mv, f"pyir_assign '{target_name}' first-def" - ) - # When the caller supplied slot context (e.g. ``self._m_idx`` - # first-def in a ``@cute.jit __init__``), register the - # freshly-created ref against the slot so a later - # ``pyir_read(... owner=..., slot_name=...)`` finds it via - # the slot registry instead of falling through to legacy - # ``_mutable_ref`` (which a wrapper rebuild can drop) and - # then to ``_create_ref`` Case-D poison fallback. - if ( - owner is not None - and slot_name is not None - and _slot_storage_available(owner) - ): - _set_slot_mv(owner, slot_name, fresh_mv) - elif _is_staged_value(new_value) and _can_create_ref(new_value): - # Outside staged CF (or non-eager type): no ref yet, but the - # local will be read inside a later staged region. Return a - # fresh wrapper so that read attaches ``_mutable_ref`` to a - # local-owned object rather than the shared source value. - log().info( - "[pyir_assign] '%s' first def → fresh wrapper (no ref yet)", - target_name, - ) - new_value = _fresh_wrapper(new_value) - else: - log().info("[pyir_assign] '%s' first def → passthrough", target_name) - # Slot-identity wrap for literal-bool / int / float first-defs inside - # staged CF. Without this, a Python ``skip = False`` (or ``count = - # 0``) first-def returns the bare primitive to the caller; the - # caller's Python local then carries no slot identity. If a later - # write inside this region promotes the slot via D1 (because the - # new value is a STAGED DSL type the keep-constexpr gate does not - # catch), the promotion lands in ``_slot_refs`` but the caller's - # local still holds the bare primitive. Downstream reads (e.g. - # ``if skip:`` after the region) then constexpr-fold with the - # stale primitive instead of going through ``pyir.load %ref`` -- - # silent miscompile. - # - # Wrapping as ``_WatchedM(value, slot_key)`` makes the slot - # identity travel with the value. After the region exits, the - # post-region merge calls ``_pyir_auto_load_arg`` on the slot's - # ``mix_iter_args`` entry; ``_WatchedM`` carrying a promoted slot - # key resolves to ``_load_as_dsl`` (which emits ``pyir.load`` and - # attaches a fresh ``_mutable_ref``) so the caller's local picks - # up the promoted SSA. Outside staged CF, ``_WatchedM`` is - # transparent (``int`` / ``float`` subclass) so consumers that - # treat it as a primitive (e.g. shape params, ``isinstance(x, - # int)``) continue to work. - # - # Skip wrapping for already-wrapped values (``_WatchedM`` re-wrap - # would lose the existing slot key) and for non-primitives (DSL - # types handled by the eager-ref branch above). - if ( - is_inside_staged_cf() - and target_name is not None - and type(new_value) in (bool, int, float) - ): - d1_slot = _make_slot_key(target_name, owner, slot_name, _current_fn_id()) - if d1_slot is not None: - # Per-iteration reset: when an earlier trace-time unrolled - # iteration of an enclosing ``range_constexpr`` already - # promoted this slot via D1, the Python-primitive first-def - # here is the user's reset at the top of the current - # unrolled body (``acc = 0.0`` / ``cur_max = -inf``). Emit - # ``pyir.store(primitive, %ref)`` at the current IP so - # downstream reads observe the reset instead of loading the - # value the previous iteration last stored, and return - # ``_load_as_dsl`` so the caller's Python local carries the - # reset SSA (a ``_WatchedM`` return would be ignored by - # ``pyir_read``'s slot-already-promoted shortcut and the - # store would race the stale load). - existing_ref = _slot_refs.get(d1_slot) - if existing_ref is not None and pyir is not None: - log().info( - "[pyir_assign] '%s' Python-primitive first-def into " - "already-promoted slot -> store + load (per-iteration reset)", - target_name, - ) - new_ir = _emit_constant_at_ip(new_value) - pyir.store(new_ir, existing_ref) - return _load_as_dsl(existing_ref, new_value) - new_value = _WatchedM(new_value, d1_slot) - # First-def depth bookkeeping for the straight-line type-change - # rebind guard (see ``_slot_first_def_depth_any``). Recorded for - # EVERY local first-def -- primitive OR staged DSL value -- so a - # later reassignment can tell whether the prior binding originated - # in the current staged-CF region (straight-line) or crosses a - # region boundary (a genuine Rule-4 join). - if target_name is not None: - _any_slot = _make_slot_key(target_name, owner, slot_name, _current_fn_id()) - if _any_slot is not None: - _slot_first_def_depth_any[_any_slot] = get_staged_cf_depth() - - # First-def location bookkeeping: record whether this slot was - # first-defined inside staged CF when the init is a Python - # primitive (bool / int / float). A later D1 reassignment - # consults this flag to refuse promotion -- placing the ref's - # init at the function-entry block would hoist the user's - # per-iteration reset out of the enclosing scf.while / scf.for. - if isinstance(new_value, (bool, int, float)) and target_name is not None: - d1_slot = _make_slot_key(target_name, owner, slot_name, _current_fn_id()) - if d1_slot is not None: - _slot_first_def_inside_cf[d1_slot] = is_inside_staged_cf() - _slot_first_def_depth[d1_slot] = get_staged_cf_depth() - try: - _slot_first_def_block[d1_slot] = ir.InsertionPoint.current.block - except Exception: - _slot_first_def_block.pop(d1_slot, None) - return new_value + return _assign_first_def( + target_name, + new_value, + filename, + lineno, + owner, + slot_name, + fresh_binding, + fresh_paths, + rhs_ctor, + ) # D1 (META_VALUE_TABLE_DESIGN): retroactive promotion of M values. # When the old value is a ``_WatchedM`` wrapper, or the slot is # already promoted, route through the meta-value table. This # handles M-mutation inside staged CF without falling through to - # the legacy "M mutation forbidden" error. - if is_inside_staged_cf() and target_name != "self.dummy": - d1_slot = _make_slot_key(target_name, owner, slot_name, _current_fn_id()) + # the "M mutation forbidden" error. + if is_inside_staged_cf(): + d1_slot = _make_slot_key(target_name, owner, slot_name) if d1_slot is not None: - # Keep-constexpr gate: refuse to promote when the slot's - # first-def was inside staged CF and the new value is still - # a Python primitive. Promotion would lift the slot's init - # to the function-entry block, hoisting the user's per- - # iteration reset out of the enclosing scf.while / scf.for. - # When the user wants persistent semantics they initialise - # the slot OUTSIDE the staged region (flag=False), which - # still promotes normally. The ``d1_slot not in _slot_refs`` - # clause keeps the rule local: once a slot is promoted (e.g. - # a sibling staged region wrote a dynamic value), continued - # Python-primitive writes go through the existing D1 path - # and store into the ref. - # - # Depth guard (L103): keep-constexpr is only sound when the - # reassignment is at the SAME staged-CF depth as the first-def - # -- i.e. the toggle constant-folds across a ``range_constexpr`` - # unroll at the same runtime-CF level (L97/L97b). When the - # current depth is GREATER, the toggle sits inside a NESTED - # RUNTIME loop/if entered after the first-def (FMHA-decode QK - # GEMM: ``scale_d`` born in an outer ``cutlass.range`` then - # toggled in a nested ``range`` MMA-K loop). That loop must - # thread the value as an iter_arg, so we must NOT keep it - # constexpr -- fall through to D1 promotion instead. - # - # Subscript exemption (L107): a SUBSCRIPT target (``"["`` in - # the name, e.g. ``_tmem_o_pw_counters[obj_id] = pw + 1``) is a - # TRACE-TIME counter incremented once per trace-visit -- never a - # loop-carried runtime value. It must stay meta regardless of - # depth (matching the dedicated subscript passthrough below). - # The depth guard alone would let a deeper-than-first-def visit - # fall through to D1 promotion (via the ``prior_meta_use`` baked - # by an earlier same-depth visit), turning the counter into a - # ``pyir.load`` -- then ``const_expr(counter % 2 == 0)`` sees a - # dynamic value. A genuine loop-carried value in a subscript - # would be a STAGED DSL type, which fails the - # ``isinstance(new_value, (bool, int, float))`` test below. + # Keep-constexpr gate: refuse to promote a still-primitive write + # whose first-def was inside staged CF (promotion would hoist the + # per-iteration reset out of the loop). + # Depth guard: keep-constexpr holds only at the SAME staged-CF + # depth as the first-def; a deeper write sits in a nested runtime + # region and must carry, so it falls through to D1 promotion. + # Subscript exemption: a subscript target is a trace-time counter + # (once per trace-visit) and stays meta regardless of depth. _is_subscript_counter = target_name is not None and "[" in target_name + # A write in a block strictly inside the first-def's block is a + # conditional update; keep-constexpr would fold the untaken branch. + # EXCEPT when every crossed region is an ``scf.if`` arm of the + # slot's own birth region (a per-iteration reset counter, e.g. + # fmha's q0_index): in-arm folds are exact per program point, so + # keep-constexpr with an escape mark (out-of-arm consumptions + # refuse at the read) is the faithful disposition. + _fd_block = _slot_first_def_block.get(d1_slot) + _reassign_strictly_nested = False + _arm_local_write = False + if _fd_block is not None and pyir is not None: + try: + _wr_block = ir.InsertionPoint.current.block + _reassign_strictly_nested = _block_strictly_inside( + _wr_block, _fd_block + ) + # Only a slot ALREADY consumed as trace-time structure + # takes the arm-local disposition: promotion cannot + # retarget its bake, so keep-constexpr + escape mark is + # the one faithful treatment. An unconsumed slot keeps + # the D1 promotion flow (a staged select carries the + # conditional update on every path). + if ( + _reassign_strictly_nested + and _PYIR_STRUCTURAL_META_CONSUMPTIONS.get(d1_slot) is not None + ): + _arm_local_write = _pyir_write_in_if_arms_of( + _wr_block, _fd_block + ) + except Exception: + _reassign_strictly_nested = False + _arm_local_write = False if ( _slot_first_def_inside_cf.get(d1_slot, False) and ( - get_staged_cf_depth() <= _slot_first_def_depth.get(d1_slot, 0) + current_staged_cf_depth() <= _slot_first_def_depth.get(d1_slot, 0) or _is_subscript_counter + or _arm_local_write ) and isinstance(new_value, (bool, int, float)) and not _is_staged_value(new_value) and isinstance(old_value, _WatchedM) and d1_slot not in _slot_refs + and (not _reassign_strictly_nested or _arm_local_write) ): + if _arm_local_write: + _PYIR_ARM_LOCAL_META_WRITES[d1_slot] = ( + ir.InsertionPoint.current.block, + filename or "", + lineno or 0, + ) + else: + # A same-depth write (the next reset) retires the mark. + _PYIR_ARM_LOCAL_META_WRITES.pop(d1_slot, None) log().info( "[pyir_assign] '%s' D1 keep-constexpr (first-def " "inside CF, new is Python primitive)", target_name, ) return _WatchedM(new_value, d1_slot) + # Past keep-constexpr the write needs promotion/carry, but a + # place consumed as trace-time structure has no retargetable + # constant: a value-changing write bakes stale structure -- refuse. + # + # Unless the write stays in the region that births the place: a + # folded ``if`` adds no staged depth, so this re-binds the place + # straight-line in the body that re-runs its binding every + # iteration. The bake therefore stays valid and the normal + # promotion flow is correct -- only a place born ABOVE the loop + # carries its value into the next iteration and is a real hazard. + _sc = _PYIR_STRUCTURAL_META_CONSUMPTIONS.get(d1_slot) + if _sc is not None and _pyir_structural_bake_is_reseeded(d1_slot): + _sc = None + if _sc is not None and _pyir_structural_value_conflicts(_sc, new_value): + raise DSLUserCodeError( + DiagId.PHASE_STRUCTURAL_CONSTANT_MUTATED, + filename=filename, + lineno=lineno, + var=target_name or str(d1_slot), + value=_sc[0], + read_file=_sc[1], + read_line=_sc[2], + ) # Gate: D1 must only fire when a PRIOR read (in an earlier # statement) baked a constant into ``_meta_uses[slot]`` -- # this is the snapshot-rewrite case D1 was designed for. @@ -774,6 +1471,7 @@ def pyir_assign( # so the baked constants from this statement's RHS get # rewritten to ``pyir.load %ref`` for correct lowering. old_is_watched = isinstance(old_value, _WatchedM) + old_is_meta_primitive = isinstance(old_value, (bool, int, float)) new_is_primitive = isinstance(new_value, (bool, int, float, _WatchedM)) new_is_staged = _is_staged_value(new_value) old_cached_ir = getattr(old_value, "_cached_ir", None) @@ -781,20 +1479,118 @@ def pyir_assign( u is not old_cached_ir for u in _meta_uses.get(d1_slot, []) ) already_tracked = d1_slot in _slot_refs or prior_meta_use_exists - # AUTO_M2S=True opt-in: when the slot has a ``_WatchedM`` old - # value (i.e. the RHS just baked a constant from this Python - # primitive) and the user explicitly enabled M->S promotion, - # also fire D1 so the baked constants get rewritten to - # ``pyir.load %ref``. Without this fall-back, the legacy M->S - # path would create a fresh ref but leave dangling constants. - # Plain dict-int subscripts (no ``_WatchedM`` wrapper) are - # never auto-promoted by this branch -- they fall through to - # the legacy passthrough so trace-time counters stay meta. - if not already_tracked and old_is_watched: - from .common import is_auto_m2s_enabled as _is_auto_m2s_enabled - - if _is_auto_m2s_enabled(): + # A same-value re-bake in the SIBLING branch of a standalone loop-free + # scf.if is a per-branch re-trace, not a carry: keep the slot meta. + _rebaked_in_sibling = False + if ( + already_tracked + and d1_slot not in _slot_refs + and not _reassign_strictly_nested + and new_is_primitive + and isinstance(old_value, _WatchedM) + and isinstance(old_value.python_value, (bool, int, float)) + and _const_values_equal(_unwrap(new_value), old_value.python_value) + and _innermost_enclosing_loop_op_at_ip() is None + ): + _baked = _meta_uses.get(d1_slot, []) + _rebaked_in_sibling = any( + _meta_use_in_sibling_if_region(_baked, _if) + for _if in _loop_free_enclosing_if_ops_at_ip() + ) + if _rebaked_in_sibling: + log().info( + "[pyir_assign] '%s' Mp->Mp re-baked constant in sibling " + "standalone-if branch -> keep meta (no auto-promote)", + target_name, + ) + _meta_uses.pop(d1_slot, None) + return _WatchedM(new_value, d1_slot) + # Meta-to-staged promotion is the DEFAULT: a meta/_WatchedM slot assigned a STAGED value inside a + # region is promoted to a function-entry ref. A plain-primitive Mp->Mp is NOT force-promoted. + # An idempotent rewrite (same primitive value AND type, e.g. a + # staged-loop body re-binding ``n = 128`` when ``n`` is already + # 128) is a no-op on every path: the baked constant stays correct + # whether or not the write executes, so it must not promote the + # slot through ANY of the Mp->Mp rules below (prior-use D1, + # gate-once, AUTO_M2S, strictly-nested). Type equality is + # required on top of ``_const_values_equal`` (which already + # splits bool from int) because ``1 == 1.0`` yet an int->float + # rebind changes the baked constant's IR type -- a real rebind. + # Each skipped write anchors its program point with a constant so + # a LATER different-value promotion can re-materialise the no-op + # as a per-site reset store (``_meta_promote_slot``); a slot that + # never promotes leaves only dead constants behind. A slot that + # already owns a cell falls through and stores through it. + if ( + d1_slot not in _slot_refs + and new_is_primitive + and isinstance(old_value, _WatchedM) + and isinstance(old_value.python_value, (bool, int, float)) + and type(_unwrap(new_value)) is type(old_value.python_value) + and _const_values_equal(_unwrap(new_value), old_value.python_value) + ): + _anchor = _emit_constant_at_current_ip(_unwrap(new_value)) + _meta_idempotent_write_anchors.setdefault(d1_slot, []).append(_anchor) + log().info( + "[pyir_assign] '%s' idempotent Mp->Mp rebind -> keep meta " + "(anchor recorded) -> %s", + target_name, + d1_slot, + ) + return _WatchedM(new_value, d1_slot) + if ( + not already_tracked + and new_is_staged + and (old_is_watched or old_is_meta_primitive) + ): + already_tracked = True + # Gate-once init (see pyir_state): the first-def was proven once-per-key, + # so a subsequent Mp->Mp write is the loop-carried advance -- promote. + if ( + not already_tracked + and old_is_watched + and new_is_primitive + and d1_slot in _PYIR_GATE_ONCE_INIT_SLOTS + ): + already_tracked = True + # Mp->Mp promotion gate: explicit AUTO_M2S promotes any slot; PyIR-implied + # promotes only a bare local or an attr leg of a container-held object. + if not already_tracked and old_is_watched and new_is_primitive: + from .common import ( + is_auto_m2s_enabled as _is_auto_m2s_enabled, + is_pyir_enabled as _is_pyir_enabled, + ) + + _is_name_slot = ( + isinstance(d1_slot, tuple) + and len(d1_slot) > 0 + and d1_slot[0] == "local" + ) or _pyir_owner_is_container_held(owner) + # A constexpr-unroll write is straight-line trace-time Python (meta), + # UNLESS the scope opened under a dynamic region and isn't idempotent. + _value_idempotent_rewrite = ( + isinstance(old_value, _WatchedM) + and isinstance(old_value.python_value, (bool, int, float)) + and _const_values_equal(_unwrap(new_value), old_value.python_value) + ) + _conditional_constexpr_write = ( + constexpr_scope_under_staged_cf() and not _value_idempotent_rewrite + ) + if _is_auto_m2s_enabled() or ( + _is_pyir_enabled() + and _is_name_slot + and (not is_inside_constexpr_loop() or _conditional_constexpr_write) + ): already_tracked = True + # Mp->Mp strictly NESTED inside the first-def block is a CONDITIONAL update (branch-dependent), + # so promote unconditionally -- without this the conditional write is dropped and the read folds. + if ( + not already_tracked + and old_is_watched + and new_is_primitive + and _reassign_strictly_nested + ): + already_tracked = True # D1 accepts: # * primitives / _WatchedM -> bake as IR constant # * staged DSL values (Int32(a+b), Boolean, ...) -> store the @@ -806,33 +1602,85 @@ def pyir_assign( initial_py = old_value.python_value if old_is_watched else old_value if isinstance(initial_py, (bool, int, float)): _meta_promote_slot( - d1_slot, initial_py, target_name, filename, lineno + d1_slot, + initial_py, + target_name, + filename, + lineno, + promoted_value=new_value, ) + # An owner place is repairable at region boundaries: + # record it for the stale-leaf reload sweep. + if d1_slot in _slot_refs and owner is not None: + _pyir_record_promoted_place_leaf(owner, slot_name) ref = _slot_refs.get(d1_slot) if ref is not None: + # A STAGED value may be trapped in a closed CF region (a post-region + # store fails dominance); the trapped-child-region guard skips it. + if new_is_staged and not isinstance(new_value, _WatchedM): + skipped = _pyir_skip_trapped_child_region_writeback( + ref, new_value, place=d1_slot + ) + if skipped is not None: + log().info( + "[pyir_assign] '%s' D1 staged value from " + "closed child region (ref already written " + "in-loop) -> skip redundant store -> %s", + target_name, + d1_slot, + ) + return skipped # Bake the new value as IR and store into the ref. - if isinstance(new_value, _WatchedM): - new_ir = new_value.ir_value() - elif new_is_staged: - new_ir = new_value.ir_value() + if isinstance(new_value, _WatchedM) or new_is_staged: + # A store captures ITS program point: the value's own + # raw backing, never an ``ir_value()`` cell re-follow. + # Exception (clause B): a region-crossed epoch-equal + # raw must observe its cell; the swap pin keeps its + # epoch-MISMATCHED capture. + from .pyir_core import _pyir_region_fresh_raw + + _fresh_rhs = ( + _pyir_region_fresh_raw(new_value) + if new_is_staged and not isinstance(new_value, _WatchedM) + else None + ) + if _fresh_rhs is not None: + new_value = _fresh_rhs + new_ir = _raw_backing_ir_value(new_value) + if new_ir is None: + new_ir = new_value.ir_value() else: - new_ir = _emit_constant_at_ip(new_value) - pyir.store(new_ir, ref) + new_ir = _emit_constant_at_current_ip(new_value) + # A rebind baked at a type other than the cell pointee is a + # fresh redefinition: re-mint the ref, republish the slot. + try: + _retyped = new_ir.type != ref.type.pointee + except Exception: + _retyped = False + if _retyped: + ref = pyir.ref(new_ir) + _slot_refs[d1_slot] = ref + # W2 type step (F-TYPEID): a staged store advances the row's + # wrapper template; a primitive store keeps the row's class + # (a re-minted generation re-books its declared class). + if new_is_staged and not isinstance(new_value, _WatchedM): + _pyir_record_slot_template(d1_slot, new_value) + elif _retyped: + _pyir_record_slot_template( + d1_slot, + _pyir_declared_promotion_template( + new_value.python_value + if isinstance(new_value, _WatchedM) + else new_value + ), + ) + _pyir_emit_store(new_ir, ref) log().info( "[pyir_assign] '%s' D1 store -> %s", target_name, d1_slot, ) - sample = ( - old_value.python_value - if old_is_watched - else ( - new_value.python_value - if isinstance(new_value, _WatchedM) - else new_value - ) - ) - return _load_as_dsl(ref, sample) + return _load_as_dsl(ref, place=d1_slot) if ( is_inside_staged_cf() @@ -841,216 +1689,1262 @@ def pyir_assign( and isinstance(new_value, (bool, int, float)) and type(old_value) is type(new_value) ): + # A dict-SUBCLASS entry reaching this meta passthrough with a CHANGED + # value was never promoted: subclasses are never adopted as watched + # owners, and the augassign spelling loads the entry raw (no read + # choke records a baked constant), so the D1 write-side promotion + # above cannot fire -- the per-iteration change would silently + # freeze at the first pass's value. Exact plain dicts keep the + # passthrough (the trace-time-counter channel); constexpr scopes + # realize the update at trace time. + if ( + isinstance(owner, dict) + and type(owner) is not dict + and not isinstance(owner, _WatchedDict) + and old_value != new_value + and not is_inside_constexpr_loop() + ): + raise DSLUserCodeError( + DiagId.CONTAINER_DICT_META_WRITE_UNPROMOTED, + filename=filename, + lineno=lineno, + var=target_name, + old_value=repr(old_value), + new_value=repr(new_value), + ) log().info( "[pyir_assign] '%s' trace-time primitive subscript update → passthrough", target_name, ) return new_value - # Straight-line type-change rebind exemption (Rule 4 false positive). - # - # A same-name local reassigned to a value of a DIFFERENT MLIR type is a - # Rule-4 "unstable join" violation ONLY at a real join -- a branch merge - # or a loop back-edge -- where the prior binding crosses a staged-CF - # boundary. But ``assign_meta_staged_check`` runs per assignment at trace - # time and cannot see CF structure, so its Rule 4 fires on EVERY - # staged->staged type-changing reassignment while ``is_inside_staged_cf()`` - # -- including legitimate straight-line rebinds where the next statement - # simply derives a new-typed value from the previous one - # (``p = base + off`` -> ``p = inttoptr(p)``). - # Non-PyIR accepts those (the Python name is just rebound). - # - # Discriminate by first-def depth: when the prior binding was first-defined - # at the SAME staged-CF depth as this reassignment, both assignments live - # in the same region (straight-line) -- not a join -- so the type change is - # safe and is materialised by the type-changing-reassignment path below (a - # fresh ref typed to the new value). When the first-def is at a SHALLOWER - # depth the prior binding pre-dates the current region and the type really - # would differ at the back-edge / merge: keep raising via Rule 4. Genuine - # region-crossing joins are additionally validated by - # ``ScfGenerator._check_region_result`` at the region boundary. - _straight_line_type_rebind = False + # Different-LENGTH tuple rebind inside staged CF is always a trace-time + # rebuild (an iter_arg cannot change shape across iterations). + + # The new tuple's staged leaves carry their own SSAs: pass it through (the + # M->M path below cannot decompose a tuple). if ( is_inside_staged_cf() - and isinstance(target_name, str) - and _is_staged_value(old_value) - and _is_staged_value(new_value) - and type(old_value) is not type(new_value) - and _staged_type_changed(old_value, new_value) + and isinstance(old_value, tuple) + and isinstance(new_value, tuple) + and len(old_value) != len(new_value) ): - _rebind_slot = _make_slot_key(target_name, owner, slot_name, _current_fn_id()) - _first_depth = ( - _slot_first_def_depth_any.get(_rebind_slot) - if _rebind_slot is not None - else None - ) - if _first_depth is not None and _first_depth >= get_staged_cf_depth(): - log().info( - "[pyir_assign] '%s' straight-line type-change rebind " - "(%s -> %s, first-def depth %d == cur depth %d) → allow", - target_name, - type(old_value).__name__, - type(new_value).__name__, - _first_depth, - get_staged_cf_depth(), + # Fail-closed: an arity-changing rebind of a LOOP-CARRIED tuple inside + # a dynamic staged loop cannot thread (an iter_arg set is arity-fixed), + # and the trace-once passthrough below would bake every consumer to the + # trace-time arity (``len()`` folds, the post-loop read serves the + # snapshot) -- a silent miscompile legacy refuses at the join as + # TYPE_UNSTABLE_JOIN (tuple<0> vs tuple<1>). Mirror that judgment + # here, where the arity mismatch is visible, with the tuple-arity + # code. Scope: only a genuine cross-iteration join refuses -- the + # binding's first-def at a SHALLOWER depth (or unrecorded: a + # param/closure capture, defined outside by construction). A + # same-region straight-line rebuild, a constexpr-governed unroll, and + # a declared-untracked slot (no place key) keep the passthrough. + if ( + not is_inside_constexpr_loop() + and _innermost_enclosing_loop_op_at_ip() is not None + ): + _arity_slot = _make_slot_key(target_name, owner, slot_name) + _arity_first_depth = ( + _slot_first_def_depth_any.get(_arity_slot) + if _arity_slot is not None + else None ) - # Re-root the variable at the new type for this region so a - # subsequent same-depth rebind is also recognised as straight-line. - _slot_first_def_depth_any[_rebind_slot] = get_staged_cf_depth() - _straight_line_type_rebind = True - # Skip Rule 4; fall through to the type-changing reassignment path. + if _arity_slot is not None and ( + _arity_first_depth is None + or _arity_first_depth < current_staged_cf_depth() + ): + raise DSLUserCodeError( + DiagId.CONTAINER_TUPLE_LENGTH_CHANGED, + filename=filename, + lineno=lineno, + var=target_name, + old=len(old_value), + new=len(new_value), + ) + log().info( + "[pyir_assign] '%s' different-length tuple rebind (%d -> %d) → " + "compile-time structural rebuild, passthrough", + target_name, + len(old_value), + len(new_value), + ) + # Generation coherence: the fresh elements store through their access + # paths' still-live cells (or mark them for a loud serve), so a later + # place-routed element read resolves this rebind, not the prior + # generation left in a read-minted row. + _pyir_route_restructured_tuple_leaves(target_name, owner, slot_name, new_value) + return new_value - # Rule checks (M mutation, type stability) — may auto-coerce scalars. - # Skipped for a straight-line type-change rebind exempted above (the - # type-changing reassignment path below materialises a fresh ref). - coerced = ( - None - if _straight_line_type_rebind - else assign_meta_staged_check( - target_name, old_value, new_value, filename, lineno + # Tuple-vs-tuple rebind inside staged CF: a tuple has no ``__dict__`` so the M->M path never + # carries its staged leaves; ``_decompose_tuple`` carries each leaf through its loop-carried ref. + if ( + is_inside_staged_cf() + and isinstance(old_value, tuple) + and isinstance(new_value, tuple) + and len(old_value) == len(new_value) + and ( + # Old tuple carries a carryable staged leaf (S->S/S->M), or the new tuple introduces one (M->S). + # A pure meta->meta tuple has no staged leaf and falls through to the M->M handling below. + any( + _is_staged_value(e) and _can_carry_leaf_ref(e) + for e in _flatten_tuple(old_value) + ) + or any( + _is_staged_value(e) and _can_carry_leaf_ref(e) + for e in _flatten_tuple(new_value) + ) ) - ) - if coerced is not None: + ): log().info( - "[pyir_assign] '%s' auto-coerced %s → %s", + "[pyir_assign] '%s' tuple-vs-tuple rebind → per-leaf decomposition", target_name, - type(new_value).__name__, - type(coerced).__name__, ) - new_value = coerced + return _decompose_tuple( + target_name, + old_value, + new_value, + filename, + lineno, + set(), + owner=owner, + slot_name=slot_name, + ) - if not is_inside_staged_cf(): - log().info("[pyir_assign] '%s' outside staged CF → passthrough", target_name) - return new_value + # Whole-DICT rebind inside staged CF: decompose per entry onto the OLD + # dict's subscript legs when key sets match and every entry is staged. - # M→M compound auto-decomposition: both old and new are compound - # objects (not directly staged) with staged leaf fields. Decompose - # into per-field pyir_assign calls so the compiler sees each SSA - # update through pyir.ref/store. + # The old dict stays the binding carrier (it accumulates the entry refs). + + # A non-decomposable rebind with staged content is REFUSED; a pure meta + # dict rebind stays a harmless trace-time replacement. The whole-rebind + # intent fact is PRODUCED by the staged-CF write funnel only: a host-phase + # rebind (no staged region open) takes the outside-CF passthrough. if ( - not _is_staged_value(new_value) - and not _is_staged_value(old_value) - and old_value is not None - and not isinstance(old_value, (int, float, bool, str, bytes, type)) - and type(old_value) is type(new_value) + is_inside_staged_cf() + and isinstance(old_value, dict) + and isinstance(new_value, dict) + and old_value is not new_value ): - if _has_decomposable_staged_fields(old_value): + _old_keys = set(dict.keys(old_value)) + _same_keys = _old_keys == set(dict.keys(new_value)) + if ( + _same_keys + and _old_keys + and all( + _is_staged_value(_e) + and _can_carry_leaf_ref(_e) + and not _is_compound_single_leaf(_e) + for _k in _old_keys + for _e in ( + dict.__getitem__(old_value, _k), + dict.__getitem__(new_value, _k), + ) + ) + ): log().info( - "[pyir_assign] '%s' M→M auto-decomposition (type=%s)", + "[pyir_assign] '%s' dict whole-rebind → per-entry " + "decomposition (%d staged entries)", target_name, - type(old_value).__name__, + len(_old_keys), ) - _decompose_m2m_assign(target_name, old_value, new_value, filename, lineno) - return old_value # I1: return old — it accumulates refs - if _has_any_staged_content(old_value): + for _k in list(dict.keys(old_value)): + # Owner-keyed pre-read FIRST: it publishes the entry's cell under + # the subscript slot (else pyir_assign would mint a TWIN cell). + _old_e = pyir_read( + f"{target_name}[{_k!r}]", + dict.__getitem__(old_value, _k), + owner=old_value, + slot_name=_k, + ) + dict.__setitem__( + old_value, + _k, + pyir_assign( + f"{target_name}[{_k!r}]", + _old_e, + dict.__getitem__(new_value, _k), + filename, + lineno, + owner=old_value, + slot_name=_k, + ), + ) + return old_value # old dict carries the entry refs + if any( + _is_staged_value(_e) + for _d in (old_value, new_value) + for _e in list(dict.values(_d)) + ): raise DSLUserCodeError( DiagId.CONTAINER_OBJECT_REPLACED, filename=filename, lineno=lineno, var=target_name, ) - # Pure meta object (no staged fields): harmless replacement. - log().info( - "[pyir_assign] '%s' pure meta compound → passthrough", + # A dict field rebuilt inside staged CF would keep its + # first-iteration entries: refuse before passing through. + _refuse_meta_compound_region_rebind( + owner, + slot_name, + old_value, + new_value, + _cur_cf_depth, target_name, + filename, + lineno, ) - return new_value - - if not _is_staged_value(new_value): log().info( - "[pyir_assign] '%s' new_value not staged → passthrough", + "[pyir_assign] '%s' pure meta dict rebind → passthrough", target_name, ) return new_value - # M->S promotion: old is meta, new is staged. - # This is only reachable when auto_m2s=True (otherwise - # assign_meta_staged_check raises above). - # Promote old_value to new_value's type, create ref at function entry. - if not _is_staged_value(old_value): - dsl_type = type(new_value) - try: - promoted = dsl_type(old_value) - except (TypeError, ValueError): - raise DSLUserCodeError( - DiagId.PHASE_CONVERSION_FAILED, - filename=filename, - lineno=lineno, - var=target_name, - new_type=dsl_type.__name__, - old_value=repr(old_value), + # Opaque-object rebind inside staged CF: ``old_value`` is an opaque handle with no staged leaves, + # an ordinary within-iteration rebind. A tracked-ref slot is loop-carried state and excluded. + # A builtin mutable ``set`` is NOT an opaque handle: it has no ``__dict__`` + # so the staged-content probe misreads it, but its rebind is meta-container + # state (``s = s | {x}``) that must reach the phase check below and refuse. + if ( + is_inside_staged_cf() + and old_value is not None + and not isinstance(old_value, (bool, int, float, str, bytes, type)) + and not isinstance(old_value, set) + and not isinstance(old_value, _WatchedM) + and ( + not _is_staged_value(old_value) + # A STAGED wrapper that owns no cellable route of its own (no + # ``_pyir_ref_supported``, leaf not scalar-carryable) is an + # opaque HANDLE for carry purposes -- e.g. an Array/Pointer + # whose single ``!llvm.ptr`` leaf the place cell carries. + or ( + not _can_create_ref(old_value) + and not _can_carry_leaf_ref(old_value) + and _is_opaque_leaf_value_tree(old_value) + ) + ) + and ( + not _has_any_staged_content(old_value) + # An all-opaque-leaf value tree is carried by opaque-leaf carry; + # ``_has_any_staged_content`` misreports its derived fields. + or _is_opaque_leaf_value_tree(old_value) + ) + ): + d1_slot = _make_slot_key(target_name, owner, slot_name) + # A place rebound to a same-type opaque value stores through its ONE + # place cell, so one carry chain spans every region level. An + # attr/subscript place cells the same way a bare local does: without + # the cell the rebind passes through and the new handle's SSA stays + # trapped in the region, which the verifier rejects at the post-join + # read. A value-protocol owner is EXCLUDED: its leaves already thread + # as region results / iter_args through the opaque-leaf region carry, + # and a second cell would carry the same leaf twice. + if ( + owner is None + or slot_name is None + or not _implements_dynamic_expression(owner) + ): + _choke_ref = _pyir_place_opaque_rebind_choke( + target_name, d1_slot, old_value, new_value + ) + if _choke_ref is not None: + return new_value + if d1_slot is None or d1_slot not in _slot_refs: + # An all-meta-primitive-field compound is NOT an opaque handle: + # a region-crossing rebind of one bakes its fields (the body is + # traced once), so it must carry per-field, not pass through. + _carrier = _pyir_carry_meta_compound_region_rebind( + old_value, + new_value, + d1_slot, + _cur_cf_depth, + target_name, + filename, + lineno, + ) + if _carrier is not None: + return _carrier + # A compound field rebuilt inside staged CF would keep its + # first-iteration content: refuse before passing through. + _refuse_meta_compound_region_rebind( + owner, + slot_name, + old_value, + new_value, + _cur_cf_depth, + target_name, + filename, + lineno, ) - if not _can_create_ref(promoted): - # Non-ref-compatible type (Pointer, Array, etc.) — passthrough. - # The non-PyIR iter_args path handles these via pytree. log().info( - "[pyir_assign] '%s' M->S but %s not ref-compatible → passthrough", + "[pyir_assign] '%s' opaque-object (%s) -> (%s) " + "within-iteration rebind → passthrough", target_name, - type(promoted).__name__, + type(old_value).__name__, + type(new_value).__name__, ) return new_value - log().info( - "[pyir_assign] '%s' M->S auto-promotion: %s(%r) -> ref", - target_name, - dsl_type.__name__, - old_value, - ) - have_slot_m2s = ( - owner is not None - and slot_name is not None - and _slot_storage_available(owner) - ) - mv: "MutableValue | None" = None - if have_slot_m2s: - existing_slot_mv = _get_slot_mv(owner, slot_name) - if existing_slot_mv is not None and existing_slot_mv._is_ref_accessible(): - mv = existing_slot_mv - if mv is None: - mv = _create_ref(promoted) - if have_slot_m2s: - _set_slot_mv(owner, slot_name, mv) - mv.store(new_value) - # Anti-aliasing: when slot context is present ALWAYS reconstruct so - # the _mutable_ref attached below cannot leak onto a shared value. - if have_slot_m2s: - new_value = mv._reconstruct(new_value.ir_value()) - else: - existing_mv = getattr(new_value, "_mutable_ref", None) - if existing_mv is not mv: - new_value = mv._reconstruct(new_value.ir_value()) - _attach_mutable_ref( - new_value, mv, f"pyir_assign '{target_name}' M->S promotion" + + # Trace-local meta-container rebind: a meta list/tuple binding + # FIRST-DEFINED in the CURRENT insertion block cannot be observed across + # a staged join (no back-edge or branch merge separates birth from + # rebind), so rebinding it -- including the list->tuple freeze whose + # leaves are staged SSAs -- is trace-time Python, exactly what non-PyIR + # executes. A binding born in an OUTER or SIBLING block keeps the + # Rule-2 refusal below. A tuple->tuple rebind stays on the existing + # per-leaf decomposition / carry routes. + if ( + is_inside_staged_cf() + and isinstance(old_value, (list, tuple)) + and isinstance(new_value, (list, tuple)) + and not _is_staged_value(old_value) + and not _is_staged_value(new_value) + and not (isinstance(old_value, tuple) and isinstance(new_value, tuple)) + ): + _mc_slot = _make_slot_key(target_name, owner, slot_name) + _mc_birth = ( + _slot_first_def_block_any.get(_mc_slot) if _mc_slot is not None else None ) - return new_value + try: + _mc_cur = ir.InsertionPoint.current.block + except Exception: + _mc_cur = None + if _mc_birth is not None and _mc_cur is not None and _mc_birth == _mc_cur: + log().info( + "[pyir_assign] '%s' trace-local meta-container rebind " + "(birth block is current) → passthrough", + target_name, + ) + return new_value - # Find or create MutableValue. - # Slot-first lookup: when owner/slot_name are supplied, the storage - # slot is authoritative (``_slot_storage_available`` is True for any - # non-None owner -- dict/list owners route through ``_slot_mvs`` - # rather than ``__dict__`` but the lookup path is the same). This - # prevents shared-value aliasing bugs where three attributes - # initialised from the same Python object collapse onto a single - # pyir.ref. + # Within-body TYPE-TRANSITION rebind: a slot last written at the SAME + # staged-CF depth rebound at an incompatible type is a fresh redefinition. + + # * S->S class change: the place re-mints at the new type -- hand back a + # fresh wrapper so a later read attaches its own cell. + + # * S->container RESTRUCTURING: the old whole-value cell is retired for the + # scope; each staged leaf's place mints its own cell on its next write. + + # A container with NO staged leaf is a genuine staged->Python demotion and + # stays on the phase-rule error path below. + if ( + old_value is not None + and new_value is not None + and _prior_binding_depth is not None + and _prior_binding_depth == _cur_cf_depth + and _binding_slot is not None + and _binding_slot not in _slot_refs + and _is_staged_value(old_value) + and type(old_value) is not type(new_value) + ): + # A place TYPE transition is an MLIR-type change; a same-MLIR-type + # wrapper-class step falls through to Rule-4 coercion → store-through. + if _is_staged_value(new_value) and _staged_type_changed(old_value, new_value): + log().info( + "[pyir_assign] '%s' within-body local type change %s -> %s " + "(prior binding depth == current depth) → fresh rebind", + target_name, + type(old_value).__name__, + type(new_value).__name__, + ) + _pyir_retire_place_row(target_name, owner, slot_name) + return _fresh_wrapper(new_value) + if isinstance(new_value, tuple) and any( + _is_staged_value(e) for e in _flatten_tuple(new_value) + ): + log().info( + "[pyir_assign] '%s' within-body restructuring rebind %s -> " + "tuple with staged leaves (prior binding depth == current " + "depth) → fresh container binding", + target_name, + type(old_value).__name__, + ) + _pyir_retire_place_row(target_name, owner, slot_name) + # Generation coherence: the fresh elements store through their + # access paths' still-live cells (or mark them for a loud serve). + _pyir_route_restructured_tuple_leaves( + target_name, owner, slot_name, new_value + ) + return new_value + + # Same-name S->S rebind widening the wrapper to a richer SUBCLASS over an + # IDENTICAL MLIR type: a fresh redefinition (no SSA join conflict can arise). + + # (Rule 4 would coerce down to the base class and the single-typed ref + # would reconstruct the base on every auto-load.) + + # Reached even when a tuple-unpack first-def leaves ``_prior_binding_depth`` + # None. A loop-carried (slot-tracked) rebind keeps its carry. + if ( + is_inside_staged_cf() + and old_value is not None + and new_value is not None + and _is_staged_value(old_value) + and _is_staged_value(new_value) + and type(old_value) is not type(new_value) + and isinstance( + new_value, type(old_value) + ) # new is a SUBCLASS of old (widening) + and not isinstance( + old_value, type(new_value) + ) # strict: old is not a subclass of new + and _mlir_types_match( + old_value, new_value + ) # identical MLIR type (no SSA-join conflict) + and ( + _binding_slot is None or _binding_slot not in _slot_refs + ) # not a carried loop carry + ): + # MLIR types match, so there is no SSA join conflict regardless of the Python wrapper class; + # avoid the spurious type-stability rejection. HOW the rebind is reconciled depends on scope: + if _ref_lives_in_strictly_enclosing_block(old_value): + # CONDITIONAL update of an OUTER-scope binding (ref minted in a + # strictly-enclosing block, rebind inside a nested staged region): + + # store back through the existing ref (identical MLIR type) so the + # lowering carries it out; return the loaded value as the binding. + _enclosing_mv = getattr(old_value, "_mutable_ref", None) + # _ref_lives_in_strictly_enclosing_block returned True, which holds + # only when a live _mutable_ref is present, so it is never None here. + assert _enclosing_mv is not None + log().info( + "[pyir_assign] '%s' conditional subclass rebind of outer-scope " + "binding %s -> %s (identical MLIR type) → store-back through ref", + target_name, + type(old_value).__name__, + type(new_value).__name__, + ) + _enclosing_mv.store(new_value) + return _enclosing_mv.load() + # FLAT same-region redefinition (ref and rebind share a block, or a tuple-unpack first-def + # left no binding): a fresh subclass wrapper lets a later read in this region mint its own ref. + log().info( + "[pyir_assign] '%s' staged rebind to richer subclass %s -> %s " + "(identical MLIR type) → fresh rebind", + target_name, + type(old_value).__name__, + type(new_value).__name__, + ) + _pyir_retire_place_row(target_name, owner, slot_name) + return _fresh_wrapper(new_value) + + # Structural tuple<->staged rebind inside staged CF: a name oscillates + # between a single staged value and a tuple/list WRAPPING it. + + # Both directions are legal whole-object replacements: every flattened + # leaf of the NEW value carries an SSA dominating the bind point. + + # Gated to a NON-carried slot so a genuine loop-carried type-changing + # carry is not silently accepted. + def _all_staged_carryable_dominating(_v: Any) -> bool: + _leaves = ( + list(_flatten_tuple(tuple(_v))) if isinstance(_v, (tuple, list)) else [_v] + ) + return bool(_leaves) and all( + _is_staged_value(_lf) + and _can_carry_leaf_ref(_lf) + and _value_dominates_current_ip(_lf) + for _lf in _leaves + ) + + def _all_staged_leaves(_v: Any) -> bool: + _leaves = ( + list(_flatten_tuple(tuple(_v))) if isinstance(_v, (tuple, list)) else [_v] + ) + return bool(_leaves) and all(_is_staged_value(_lf) for _lf in _leaves) + + _old_is_tuple = isinstance(old_value, (tuple, list)) + _new_is_tuple = isinstance(new_value, (tuple, list)) + if ( + is_inside_staged_cf() + and old_value is not None + # Exactly ONE side is a tuple -- a structural wrap/unwrap, not a same-shape carry (a tuple->tuple + # carry is owned by the tuple-vs-tuple decompose gate above; a staged->staged carry by Rule 4). + and (_old_is_tuple != _new_is_tuple) + # OLD is already tracked through its leaves (single staged value, or tuple of all-staged leaves). + and _all_staged_leaves(old_value) + # NEW's every flattened leaf is a self-dominating carryable staged SSA. + and _all_staged_carryable_dominating(new_value) + ): + _stl_slot = _make_slot_key(target_name, owner, slot_name) + if _stl_slot is None or _stl_slot not in _slot_refs: + log().info( + "[pyir_assign] '%s' structural tuple<->staged rebind " + "(old leaves staged, new leaves self-dominate) → passthrough " + "(region/loop carry owns the carry)", + target_name, + ) + # Re-tupling direction: the fresh elements store through their + # access paths' still-live cells (generation coherence). + if _new_is_tuple: + _pyir_route_restructured_tuple_leaves( + target_name, owner, slot_name, new_value + ) + return new_value + + # Straight-line type-change rebind exemption (Rule 4 false positive). # - # When slot lookup returns nothing we also consult the value's - # ``_mutable_ref`` so a ref created earlier via the legacy path - # (e.g. an ``attach_ref=False`` read) is adopted by the slot instead - # of being duplicated. + # A same-name local reassigned to a value of a DIFFERENT MLIR type is a + # Rule-4 "unstable join" violation ONLY at a real join -- a branch merge + # or a loop back-edge -- where the prior binding crosses a staged-CF + # boundary. But ``assign_meta_staged_check`` runs per assignment at trace + # time and cannot see CF structure, so its Rule 4 fires on EVERY + # staged->staged type-changing reassignment while ``is_inside_staged_cf()`` + # -- including legitimate straight-line rebinds where the next statement + # simply derives a new-typed value from the previous one + # (``p = base + off`` -> ``p = inttoptr(p)``). + # Non-PyIR accepts those (the Python name is just rebound). + # + # Discriminate by first-def depth: when the prior binding was first-defined + # at the SAME staged-CF depth as this reassignment, both assignments live + # in the same region (straight-line) -- not a join -- so the type change is + # safe and is materialised by the type-changing-reassignment path below (a + # fresh ref typed to the new value). When the first-def is at a SHALLOWER + # depth the prior binding pre-dates the current region and the type really + # would differ at the back-edge / merge: keep raising via Rule 4. Genuine + # region-crossing joins are additionally validated by + # ``ScfGenerator._check_region_result`` at the region boundary. + _straight_line_type_rebind = False + if ( + is_inside_staged_cf() + and isinstance(target_name, str) + and _is_staged_value(old_value) + and _is_staged_value(new_value) + and type(old_value) is not type(new_value) + and _staged_type_changed(old_value, new_value) + ): + _rebind_slot = _make_slot_key(target_name, owner, slot_name) + _first_depth = ( + _slot_first_def_depth_any.get(_rebind_slot) + if _rebind_slot is not None + else None + ) + if _first_depth is not None and _first_depth >= current_staged_cf_depth(): + log().info( + "[pyir_assign] '%s' straight-line type-change rebind " + "(%s -> %s, first-def depth %d == cur depth %d) → allow", + target_name, + type(old_value).__name__, + type(new_value).__name__, + _first_depth, + current_staged_cf_depth(), + ) + # Re-root the variable at the new type for this region so a + # subsequent same-depth rebind is also recognised as straight-line. + _slot_first_def_depth_any[_rebind_slot] = current_staged_cf_depth() + _straight_line_type_rebind = True + # Skip Rule 4; fall through to the type-changing reassignment path. + + # Rule checks (M mutation, type stability) — may auto-coerce scalars. + # Skipped for a straight-line type-change rebind exempted above (the + # type-changing reassignment path below materialises a fresh ref). + coerced = ( + None + if _straight_line_type_rebind + else assign_meta_staged_check( + target_name, old_value, new_value, filename, lineno, owner=owner + ) + ) + if coerced is not None: + log().info( + "[pyir_assign] '%s' auto-coerced %s → %s", + target_name, + type(new_value).__name__, + type(coerced).__name__, + ) + new_value = coerced + + if not is_inside_staged_cf(): + return _assign_outside_cf( + target_name, old_value, new_value, filename, lineno, owner, slot_name + ) + + if ( + not _is_staged_value(new_value) + and not _is_staged_value(old_value) + and old_value is not None + and not isinstance(old_value, (int, float, bool, str, bytes, type)) + and ( + type(old_value) is type(new_value) + # An adopted watched container is the same PLACE as the plain container + # literal that rebinds it -- adoption must not change the dispatch. + or (isinstance(old_value, _WatchedList) and type(new_value) is list) + or (isinstance(old_value, _WatchedDict) and type(new_value) is dict) + ) + ): + return _assign_m2m_compound( + target_name, + old_value, + new_value, + filename, + lineno, + owner, + slot_name, + _binding_slot, + _cur_cf_depth, + ) + + if not _is_staged_value(new_value): + log().info( + "[pyir_assign] '%s' new_value not staged → passthrough", + target_name, + ) + return new_value + + if not _is_staged_value(old_value): + return _assign_m2s_promotion( + target_name, old_value, new_value, filename, lineno, owner, slot_name + ) + + return _assign_store_through( + target_name, old_value, new_value, filename, lineno, owner, slot_name + ) + + +# First-time definition — create ref eagerly if inside staged CF. +# Only create refs for scalar-like types (Numeric, not Boolean). +# Multi-element types (Vector), boolean types (i1), and types with +# complex ir_value() are left to the reassignment path to handle. +# +# Anti-aliasing: ``a = b = c = seed`` binds three Python locals to +# the same value object. ``ast.Name`` targets carry no storage +# owner, so ``pyir_assign``/``pyir_read`` fall back to the +# value-keyed ``_mutable_ref`` cache. Without a fresh wrapper per +# first-def, the three locals would alias onto ``seed._mutable_ref`` +# and collapse onto a single ``pyir.ref``. Returning a fresh +# wrapper per first-def gives each local its own attachment slot. +def _assign_first_def( + target_name: Any, + new_value: Any, + filename: "str | None", + lineno: "int | None", + owner: Any, + slot_name: Any, + fresh_binding: bool, + fresh_paths: tuple, + rhs_ctor: Any, +) -> Any: + # Freeze a fresh plain-local binding wired to a foreign storage slot, or + # the local re-loads the source at its consumer (breaks value semantics). + + # Fires at ANY scope: Python binds the value at assignment time whether + # or not staged CF is open. + if owner is None and slot_name is None and _carries_foreign_slot_binding(new_value): + new_value = _freeze_foreign_slot_binding( + new_value, + f"first-def local '{target_name}'", + fresh_container=fresh_binding, + ) + if owner is None and slot_name is None and isinstance(new_value, tuple): + new_value = _unalias_tuple_leaves(new_value) + # Fresh-container first-def inside staged CF re-executes per iteration: + # every decomposable staged entry re-inits its cell (declared fresh paths only). + if ( + is_inside_staged_cf() + and fresh_binding + and isinstance(new_value, (_WatchedDict, _WatchedList)) + ): + _pyir_fresh_container_init( + str(target_name), + new_value, + filename, + lineno, + frozenset(tuple(_p) for _p in fresh_paths), + ) + # The default first-def gate excludes booleans because a free- + # standing ``cond = i > 2`` is almost always a one-shot if-test + # and the dead ref breaks PYIRToSCF. But when the caller supplied + # slot context (e.g. ``self._is_valid_tile = Boolean(valid)`` in + # a jit ``__init__``), the slot store records the ref so + # downstream ``pyir_read('container', container)`` can refresh it + # across staged CF boundaries. Slot-keyed booleans are therefore + # allowed to take refs. + # + # CRITICAL: keep the gate as a single ``and`` chain so the early + # checks short-circuit BEFORE we call ``_is_boolean_like`` / + # ``_is_vector_like``. Those helpers call ``value.ir_value()`` + # which materialises (and caches) the leaf ``arith.constant`` for + # ``_WatchedInt`` -- if we evaluate them on a meta primitive we + # pin the constant at the wrong insertion point and break the + # sibling-region constant cache. + if ( + is_inside_staged_cf() + and _is_staged_value(new_value) + and _can_create_ref(new_value) + and ( + # A live row at the local place admits the write regardless of + # the value kind: every choked write of a row-backed place must + # store through the row (INV-1; vectors included). + ( + owner is None + and slot_name is None + and _get_slot_mv(None, target_name) is not None + ) + or ( + not _is_vector_like(new_value) + and ( + # Boolean admitted with an explicit slot identity, or + # inside a staged loop body (the birth-block seed store + # keeps the ref live). + not _is_boolean_like(new_value) + or bool(_pyir_open_loop_body_blocks) + or ( + owner is not None + and slot_name is not None + and _slot_storage_available(owner) + ) + ) + ) + ) + ): + log().info( + "[pyir_assign] '%s' first def inside staged CF → create ref", + target_name, + ) + new_value = _fresh_wrapper(new_value) + # F-BIRTHPOS: the first-def choke runs at the binding position, so + # the current block IS the place's birth block. + try: + _birth_block = ir.InsertionPoint.current.block + except Exception: + _birth_block = None + fresh_mv = _create_ref(new_value, birth_block=_birth_block) + if owner is None and slot_name is None: + # F-PLACE: a bare-local first-def cell is the local place's + # row; registration stamps the route a later read validates + # (V-1) and converges sibling-region first-defs on one cell. + fresh_mv = _set_slot_mv(None, target_name, fresh_mv) + fresh_mv.store(new_value) + _attach_mutable_ref( + new_value, fresh_mv, f"pyir_assign '{target_name}' first-def" + ) + # When the caller supplied slot context (e.g. ``self._m_idx`` + # first-def in a jit ``__init__``), register the + # freshly-created ref against the slot so a later + # ``pyir_read(... owner=..., slot_name=...)`` finds it via + # the slot registry instead of falling through to the value-keyed + # ``_mutable_ref`` (which a wrapper rebuild can drop) and + # then to ``_create_ref`` Case-D poison fallback. + if ( + owner is not None + and slot_name is not None + and _slot_storage_available(owner) + ): + _canon_mv = _set_slot_mv(owner, slot_name, fresh_mv) + if _canon_mv is fresh_mv: + # Fresh in-region attr first-def (no adopted prior cell): + # record it so an uncovered read refuses at trace time. + _pyir_record_cf_attr_first_def(owner, slot_name) + elif _is_staged_value(new_value) and _can_create_ref(new_value): + # Outside staged CF (or non-eager type): no ref yet, but a later + # staged read must attach ``_mutable_ref`` to a local-owned object. + + # EXCEPT a no-op op returning self (e.g. a same-dtype cast): a + # re-wrap would rebuild the object and break Python ``is`` identity. + + # Guarded to compounds (``_shape``) so SSA-backed scalars keep + # their distinct-outer-wrapper behaviour. + if ( + getattr(new_value, "_pyir_local_wrapper", False) + and getattr(new_value, "_shape", None) is not None + ): + log().info( + "[pyir_assign] '%s' first def → no-op self-return, " + "preserve identity (skip re-wrap)", + target_name, + ) + else: + log().info( + "[pyir_assign] '%s' first def → fresh wrapper (no ref yet)", + target_name, + ) + new_value = _fresh_wrapper(new_value) + else: + log().info("[pyir_assign] '%s' first def → passthrough", target_name) + # Fresh-object first-def inside staged CF (witnessed direct-ctor RHS): + # record per-field first-def facts so promotion resets at this block. + if is_inside_staged_cf() and rhs_ctor is not None: + _pyir_record_fresh_object_leaf_first_defs(rhs_ctor, new_value) + # First-def depth bookkeeping for the straight-line type-change + # rebind guard (see ``_slot_first_def_depth_any``). Recorded for + # EVERY local first-def -- primitive OR staged DSL value -- so a + # later reassignment can tell whether the prior binding originated + # in the current staged-CF region (straight-line) or crosses a + # region boundary (a genuine Rule-4 join). + if target_name is not None: + _any_slot = _make_slot_key(target_name, owner, slot_name) + if _any_slot is not None: + _slot_first_def_depth_any[_any_slot] = current_staged_cf_depth() + # F-BIRTHPOS: the binding's birth block, for every first-def + # (a later same-block rebind is provably trace-local). + try: + _slot_first_def_block_any[_any_slot] = ir.InsertionPoint.current.block + except Exception: + _slot_first_def_block_any.pop(_any_slot, None) + + # Slot-identity wrap + per-iteration reset + first-def location + # bookkeeping for literal first-defs; tuples recurse per leaf. + if target_name is not None: + new_value = _record_meta_primitive_first_def( + target_name, new_value, owner, slot_name + ) + return new_value + + +def _assign_outside_cf( + target_name: Any, + old_value: Any, + new_value: Any, + filename: "str | None", + lineno: "int | None", + owner: Any, + slot_name: Any, +) -> Any: + log().info("[pyir_assign] '%s' outside staged CF → passthrough", target_name) + # An already-PROMOTED bare-local slot's cell stays authoritative even + # for a Python-phase rebind: store through it so later reads re-load. + if owner is None: + _pt_slot = _make_slot_key(target_name, owner, slot_name) + _pt_ref = _slot_refs.get(_pt_slot) if _pt_slot is not None else None + # OPAQUE-row generation hygiene: a Python-phase rebind of a local + # with a dialect-opaque cell stores through or retires the row. + if _pt_ref is not None and pyir is not None: + try: + _pt_pointee = _pt_ref.type.pointee + except Exception: + _pt_pointee = None + if _pt_pointee is not None and _pyir_ir_type_is_opaque(_pt_pointee): + _pt_new_raw = _raw_backing_ir_value(new_value) + if ( + isinstance(_pt_new_raw, ir.Value) + and _pt_new_raw.type == _pt_pointee + and _value_dominates_current_ip(_pt_ref) + ): + _pyir_emit_store( + _pt_new_raw, + _pt_ref, + choke=f"Python-phase opaque rebind '{target_name}'", + ) + _pyir_record_slot_template(_pt_slot, new_value) + else: + _slot_refs.pop(_pt_slot, None) + _slot_templates.pop(_pt_slot, None) + return new_value + if _pt_ref is not None and pyir is not None: + _pt_is_primitive = isinstance(new_value, (bool, int, float, _WatchedM)) + _pt_is_staged = _is_staged_value(new_value) and not ( + isinstance(new_value, _WatchedM) + ) + if _pt_is_staged and not _is_compound_single_leaf(new_value): + _pt_ir = new_value.ir_value() + if _pyir_ref_pointee_type_changed(_pt_ref, new_value): + # Fresh redefinition at a new width: re-mint and + # republish, exactly like the in-CF D1 store path. + _pt_ref = pyir.ref(_pt_ir) + _slot_refs[_pt_slot] = _pt_ref + # F-TYPEID: the store advances the row's wrapper template. + _pyir_record_slot_template(_pt_slot, new_value) + _pyir_emit_store(_pt_ir, _pt_ref) + log().info( + "[pyir_assign] '%s' Python-phase rebind of promoted " + "slot → store-through (cell stays authoritative)", + target_name, + ) + return _load_as_dsl(_pt_ref, place=_pt_slot) + if _pt_is_primitive: + _pt_py = ( + new_value.python_value + if isinstance(new_value, _WatchedM) + else new_value + ) + _pyir_emit_store(_emit_constant_for_ref(_pt_ref, _pt_py), _pt_ref) + log().info( + "[pyir_assign] '%s' Python-phase primitive rebind of " + "promoted slot → store-through", + target_name, + ) + return _load_as_dsl(_pt_ref, place=_pt_slot) + # A VALUE-KEYED place row (a local promoted by a prior staged + # region, never published as a D1 `_slot_refs` row) is equally + # authoritative for later reads (R1): a Python-phase rebind must + # store through it, or the next staged region reads the stale + # pre-rebind cell. + if _pt_ref is None and pyir is not None: + _pt_mv = _get_slot_mv(None, target_name) + if ( + _pt_mv is not None + and _pt_mv._is_ref_accessible() + and _is_staged_value(new_value) + and not isinstance(new_value, _WatchedM) + and _can_carry_leaf_ref(new_value) + and not _is_compound_single_leaf(new_value) + # A type-changing rebind of the promoted row (an + # Int32-promoted slot rebound to a Boolean, say) must not + # store through: the write is type-unfaithful and the cell + # cannot be re-minted under a live while-carry. Fall through + # so the join check refuses it as TYPE_UNSTABLE_JOIN. + and _mlir_types_match(_pt_mv._value, new_value) + ): + # INV-1' no-op: the binding already IS the row's content. + if ( + getattr(new_value, "_mutable_ref", None) is _pt_mv + and new_value is _pt_mv._value + ): + return new_value + # The guard above proved the pointee is unchanged, so this + # store is faithful (a same-type re-init of the row). + _pt_mv.store(new_value) + new_value = _pt_mv._reconstruct(new_value.ir_value()) + _attach_mutable_ref( + new_value, + _pt_mv, + f"pyir_assign '{target_name}' Python-phase row rebind", + ) + _pyir_adopt_stored_representative(_pt_mv, new_value) + log().info( + "[pyir_assign] '%s' Python-phase rebind of promoted " + "place row → store-through (row stays authoritative)", + target_name, + ) + return new_value + # Copy-on-bind alias split: a Python-phase rebind re-keys a _WatchedM to the + # target slot / fresh-wraps a staged scalar so two names never share a ref. + if new_value is not old_value: + if isinstance(new_value, _WatchedM): + _tgt_slot = _make_slot_key(target_name, owner, slot_name) + if _tgt_slot is not None and new_value._slot_key != _tgt_slot: + log().info( + "[pyir_assign] '%s' copy-on-bind split: _WatchedM " + "re-keyed %s -> %s", + target_name, + new_value._slot_key, + _tgt_slot, + ) + new_value = _WatchedM(new_value.python_value, _tgt_slot) + elif ( + _is_staged_value(new_value) + and _can_carry_leaf_ref(new_value) + and not _is_compound_single_leaf(new_value) + ): + _split = _fresh_wrapper(new_value) + if _split is not new_value: + log().info( + "[pyir_assign] '%s' copy-on-bind split: fresh scalar wrapper (%s)", + target_name, + type(new_value).__name__, + ) + new_value = _split + # Python-phase polarity: the NEW object becomes the binding; stamp the + # OLD generation superseded with pre-adoption store versions. + if ( + old_value is not None + and old_value is not new_value + and type(old_value) is type(new_value) + and _pyir_is_generation_compound(old_value) + ): + try: + _rebind_slot = _make_slot_key(target_name, owner, slot_name) + except Exception: + _rebind_slot = None + _pyir_record_generation_rebind( + _rebind_slot, + old_value, + new_value, + target_name, + filename, + lineno, + superseded_obj=old_value, + ) + # A Python-phase rebind still lands on places: leaves whose place owns + # a live cell get the store + re-alias so later reads resolve ONE cell. + try: + _pyir_route_binding_to_live_place_cells(owner, slot_name, new_value) + except Exception: + pass + return new_value + + +# M→M compound auto-decomposition: both old and new are compound +# objects (not directly staged) with staged leaf fields. Decompose +# into per-field pyir_assign calls so the compiler sees each SSA +# update through pyir.ref/store. +def _assign_m2m_compound( + target_name: Any, + old_value: Any, + new_value: Any, + filename: "str | None", + lineno: "int | None", + owner: Any, + slot_name: Any, + _binding_slot: Any, + _cur_cf_depth: int, +) -> Any: + # A value-tree whole-object BRANCH-SELECT must NOT take the ``__dict__`` + # walk (it would store a host-computed sub-leaf into a kernel ref). + + # Passing through lets the region-close ``scf.if`` carry only the + # differing canonical leaves; the reconstruct restores prototype fields. + + # Both the differing-prototype arm and the no-op arm (``old is new``) + # pass through, or the no-op arm re-creates the bad host-value stores. + + # Gated to a LOOP-FREE ``scf.if``: a per-iteration meta change must + # still hit the meta-primitive guard via the ``__dict__`` walk. + if ( + _implements_dynamic_expression(old_value) + and ( + old_value is new_value + or _value_trees_different_prototype(old_value, new_value) + ) + and _innermost_enclosing_loop_op_at_ip() is None + and _loop_free_enclosing_if_ops_at_ip() + ): + log().info( + "[pyir_assign] '%s' value-tree whole-object branch-select rebind " + "→ passthrough (region-close carry owns the carry)", + target_name, + ) + return new_value + # Whole-object rebind of a top-level LIST of ref-carryable staged + # scalars: per-element ``pyir_assign`` on the OLD list's subscript slots. + + # The old list stays the binding carrier (it accumulates element refs). + + # Gated to ALL elements ref-carryable staged on both sides; otherwise + # fall through to the existing decompose / rejection logic unchanged. + + # ``old_value is new_value`` takes the same path: each element + # self-assign stores the current SSA through the element's slot ref. + if ( + isinstance(old_value, list) + and old_value + and len(old_value) == len(new_value) + and all( + _is_staged_value(e) + and _can_carry_leaf_ref(e) + and not _is_compound_single_leaf(e) + for pair in zip(old_value, new_value) + for e in pair + ) + ): + log().info( + "[pyir_assign] '%s' list whole-rebind → per-element " + "decomposition (%d staged elements)", + target_name, + len(old_value), + ) + for _i, (_old_elem, _new_elem) in enumerate(zip(old_value, new_value)): + old_value[_i] = pyir_assign( + f"{target_name}[{_i}]", + _old_elem, + _new_elem, + filename, + lineno, + owner=old_value, + slot_name=_i, + ) + return old_value # old list carries the element refs + if _has_decomposable_staged_fields(old_value): + log().info( + "[pyir_assign] '%s' M→M auto-decomposition (type=%s)", + target_name, + type(old_value).__name__, + ) + # m2m polarity: the binding keeps the old object as carrier; the + # generation bump lets a pre-rebind alias capture detect staleness. + _gen_rebind_rec = _pyir_record_generation_rebind( + _binding_slot, + old_value, + new_value, + target_name, + filename, + lineno, + allow_empty_cells=True, + ) + _PYIR_M2M_DECOMPOSE_DEPTH[0] += 1 + try: + with _PyirGenEventSuppress(): + _decompose_m2m_assign( + target_name, old_value, new_value, filename, lineno + ) + finally: + _PYIR_M2M_DECOMPOSE_DEPTH[0] -= 1 + # Cells the decompose walk itself minted enter the baseline as + # born-at-rebind (a pre-rebind alias read of them always diverges). + _pyir_complete_generation_rebind_cells(_gen_rebind_rec, old_value) + return old_value # I1: return old — it accumulates refs + if _has_any_staged_content(old_value): + # An opaque-leaf value-tree is not scalar-decomposable but is legal: + # its carry carries on its canonical extract-leaf at region-close. + if _is_opaque_leaf_value_tree(old_value): + log().info( + "[pyir_assign] '%s' opaque-leaf value-tree rebind → " + "passthrough (region-close carry owns the carry)", + target_name, + ) + return new_value + # MIXED-leaf value-tree generalisation: a fresh same-class rebind + # whose NEW canonical leaves all DOMINATE the bind point loses no carry. + + # Opaque leaves carry at region-close, scalar leaves via the M2S + # slot machinery -- pass through, as non-PyIR rebinds the local. + + # Gated to a NON-carried slot and all-dominating NEW leaves so a + # genuine region-crossing type-changing carry is not accepted. + _co_slot = _make_slot_key(target_name, owner, slot_name) + if _co_slot is None or _co_slot not in _slot_refs: + _new_leaves = _pyir_extract_leaf_values(new_value) + if _new_leaves and all( + _value_dominates_current_ip(_lf) for _lf in _new_leaves + ): + log().info( + "[pyir_assign] '%s' mixed-leaf value-tree whole-object rebind " + "(all new canonical leaves self-dominate) → passthrough " + "(region-close / M2S carry owns the carry)", + target_name, + ) + return new_value + raise DSLUserCodeError( + DiagId.CONTAINER_OBJECT_REPLACED, + filename=filename, + lineno=lineno, + var=target_name, + ) + # Pure meta object (no staged fields): harmless replacement, except + # for a region-crossing all-meta-primitive rebind, which carries + # per-field (the old object stays the binding carrier). + _carrier = _pyir_carry_meta_compound_region_rebind( + old_value, + new_value, + _binding_slot, + _cur_cf_depth, + target_name, + filename, + lineno, + ) + if _carrier is not None: + return _carrier + log().info( + "[pyir_assign] '%s' pure meta compound → passthrough", + target_name, + ) + return new_value + + +# M->S promotion: old is meta, new is staged. +# This is only reachable when auto_m2s=True (otherwise +# assign_meta_staged_check raises above). +# Promote old_value to new_value's type, create ref at function entry. +def _assign_m2s_promotion( + target_name: Any, + old_value: Any, + new_value: Any, + filename: "str | None", + lineno: "int | None", + owner: Any, + slot_name: Any, +) -> Any: + dsl_type = type(new_value) + try: + promoted = dsl_type(old_value) + except (TypeError, ValueError): + raise DSLUserCodeError( + DiagId.PHASE_CONVERSION_FAILED, + filename=filename, + lineno=lineno, + var=target_name, + new_type=dsl_type.__name__, + old_value=repr(old_value), + ) + if not _can_create_ref(promoted): + # Non-ref-compatible type (Pointer, Array, etc.) — passthrough. + # The non-PyIR iter_args path handles these via pytree. + log().info( + "[pyir_assign] '%s' M->S but %s not ref-compatible → passthrough", + target_name, + type(promoted).__name__, + ) + return new_value + log().info( + "[pyir_assign] '%s' M->S auto-promotion: %s(%r) -> ref", + target_name, + dsl_type.__name__, + old_value, + ) + have_slot_m2s = ( + owner is not None and slot_name is not None and _slot_storage_available(owner) + ) + mv: "MutableValue | None" = None + if have_slot_m2s: + existing_slot_mv = _get_slot_mv(owner, slot_name) + if existing_slot_mv is not None and existing_slot_mv._is_ref_accessible(): + mv = existing_slot_mv + if mv is None: + mv = _create_ref( + promoted, + birth_block=_pyir_recorded_birth_block(target_name, owner, slot_name), + ) + if have_slot_m2s: + _set_slot_mv(owner, slot_name, mv) + elif owner is None and slot_name is None: + # F-PLACE: a bare-local promotion cell is the local place's + # row; registration stamps the route a later read validates. + mv = _set_slot_mv(None, target_name, mv) + mv.store(new_value) + # Anti-aliasing: when slot context is present ALWAYS reconstruct so + # the _mutable_ref attached below cannot leak onto a shared value. + if have_slot_m2s: + new_value = mv._reconstruct(new_value.ir_value()) + else: + existing_mv = getattr(new_value, "_mutable_ref", None) + if existing_mv is not mv: + new_value = mv._reconstruct(new_value.ir_value()) + _attach_mutable_ref(new_value, mv, f"pyir_assign '{target_name}' M->S promotion") + # The reconstructed wrapper rewraps the just-stored raw: record it as + # the cell's canonical representative for the staged-read choke. + _pyir_adopt_stored_representative(mv, new_value) + return new_value + + +def _assign_store_through( + target_name: Any, + old_value: Any, + new_value: Any, + filename: "str | None", + lineno: "int | None", + owner: Any, + slot_name: Any, +) -> Any: + # Find or create MutableValue. have_slot = ( owner is not None and slot_name is not None and _slot_storage_available(owner) ) mv = None if have_slot: mv = _get_slot_mv(owner, slot_name) - if mv is None: - legacy_mv = getattr(old_value, "_mutable_ref", None) - if legacy_mv is not None: - mv = legacy_mv - _set_slot_mv(owner, slot_name, mv) else: - mv = getattr(old_value, "_mutable_ref", None) + # INV-1': the resolved place row is the commit target whenever one + # exists -- the old value's route is a fact about the VALUE, never + # about the row's currency, so route equality can never dedup a + # commit the row has not seen. + if owner is None and slot_name is None: + mv = _get_slot_mv(None, target_name) + if mv is None: + mv = getattr(old_value, "_mutable_ref", None) log().info("[pyir_assign] '%s' existing mv=%s", target_name, mv) + # The type-change predicate is pure in (old_value, new_value) and neither + # is rebound below, so its three sites here share one evaluation. + _stc_memo: "bool | None" = None + + def _staged_type_changed_once() -> bool: + nonlocal _stc_memo + if _stc_memo is None: + _stc_memo = _staged_type_changed(old_value, new_value) + return _stc_memo + # Type-changing reassignment: a same-name local rebound to a value of # a DIFFERENT MLIR type (e.g. ``v = v.to(other_dtype)`` or a vector # recast ``vector<8xf4E2M1FN> -> vector<4xi8>``) cannot reuse the old @@ -1061,9 +2955,9 @@ def pyir_assign( # value so the ref/store/load all carry the new type. ``_can_create_ # ref(new_value)`` gates this so non-ref types still passthrough. # Layer-4 guard (P89): a SAME-Python-type, ref-supported staged value - # (e.g. cute._Tensor -> cute._Tensor) whose MLIR type CHANGES at a region + # (same wrapper class on both sides) whose MLIR type CHANGES at a region # crossing must NOT silently create a fresh, non-escaping ref -- the new - # value would be computed but never threaded out (no scf result / iter_arg), + # value would be computed but never carried out (no scf result / iter_arg), # and the post-CF use would read the STALE pre-CF value (a silent # miscompile). Rule 4 (assign_meta_staged_check) misses this because it only # compares the Python class, and both sides are the same class here. Honor @@ -1077,18 +2971,18 @@ def pyir_assign( and type(old_value) is type(new_value) and _can_create_ref(old_value) and _can_create_ref(new_value) - and _staged_type_changed(old_value, new_value) + and _staged_type_changed_once() ): - _tc_slot = _make_slot_key(target_name, owner, slot_name, _current_fn_id()) + _tc_slot = _make_slot_key(target_name, owner, slot_name) _tc_first_depth = ( _slot_first_def_depth_any.get(_tc_slot) if _tc_slot is not None else None ) _is_straight_line = ( - _tc_first_depth is not None and _tc_first_depth >= get_staged_cf_depth() + _tc_first_depth is not None and _tc_first_depth >= current_staged_cf_depth() ) if not _is_straight_line: # Surface the old vs new value in its human-readable form -- the - # same repr ``print(x)`` shows (e.g. a cute._Tensor prints as + # same repr ``print(x)`` shows (e.g. a tensor wrapper prints as # ``tensor> o (128,128):(1,128)>``), so the # user can see exactly what changed. def _type_desc(v: Any) -> str: @@ -1109,9 +3003,7 @@ def _type_desc(v: Any) -> str: ) type_changed = ( - mv is not None - and _can_create_ref(new_value) - and _staged_type_changed(old_value, new_value) + mv is not None and _can_create_ref(new_value) and _staged_type_changed_once() ) if type_changed: log().info( @@ -1132,7 +3024,7 @@ def _type_desc(v: Any) -> str: elif ( mv is not None and not _can_create_ref(new_value) - and _staged_type_changed(old_value, new_value) + and _staged_type_changed_once() ): log().info( "[pyir_assign] '%s' type-changing rebind to non-ref type " @@ -1145,9 +3037,10 @@ def _type_desc(v: Any) -> str: _clear_slot_mv(owner, slot_name) return new_value - if mv is not None and mv._is_ref_accessible(): + reuse_accessible_ref = mv is not None and mv._is_ref_accessible() + if reuse_accessible_ref: log().info("[pyir_assign] '%s' reuse ref (accessible)", target_name) - elif not _can_create_ref(old_value) and not type_changed: + elif not _can_carry_leaf_ref(old_value) and not type_changed: # Non-ref-compatible type (Pointer, Array, etc.) — passthrough. # The non-PyIR iter_args path handles these via pytree. log().info( @@ -1160,11 +3053,57 @@ def _type_desc(v: Any) -> str: # On a type-changing rebind the ref must be typed to ``new_value`` # (``old_value``'s type no longer fits). For same-type rebinds keep # the historical behaviour and build from ``old_value``. - mv = _create_ref(new_value if type_changed else old_value) + # + # F-BIRTHPOS: a type change births a fresh generation HERE; a same-type + # lazy mint seeds at the place's RECORDED first-def block (if any). + if type_changed: + try: + _birth_block = ir.InsertionPoint.current.block + except Exception: + _birth_block = None + else: + _birth_block = _pyir_recorded_birth_block(target_name, owner, slot_name) + mv = _create_ref( + new_value if type_changed else old_value, birth_block=_birth_block + ) if have_slot: _set_slot_mv(owner, slot_name, mv) + else: + if owner is None and slot_name is None: + # F-PLACE: a bare-local rebind cell is the local place's row; + # registration stamps the route a later read validates (V-1). + mv = _set_slot_mv(None, target_name, mv) + # Publish a fresh, function-entry-dominating ref under the slot's D1 + # name key so a post-region read loads it (conditional-def carry). + try: + d1_slot = _make_slot_key(target_name, owner, slot_name) + if ( + d1_slot is not None + and d1_slot not in _slot_refs + and _ref_dominates_whole_function(mv) + ): + _slot_refs[d1_slot] = mv.ref + _pyir_record_slot_template(d1_slot, mv._value) + log().info( + "[pyir_assign] '%s' published entry-block ref as D1 slot", + target_name, + ) + except Exception: + pass log().info("[pyir_assign] '%s' ref created via _create_ref", target_name) + # Every non-passthrough path above bound a cell: reuse or fresh mint. + assert mv is not None + + # INV-1' no-op re-walk: an unchanged binding's cell IS the row cell, so a + # re-observation walk commits nothing (zero new write positions). + if ( + reuse_accessible_ref + and getattr(new_value, "_mutable_ref", None) is mv + and new_value is mv._value + ): + return new_value + # P-058-D: Auto-load new_value when it carries a ref from a different # context (cross-dict propagation in copy_consumer_vars_to). # For Name/Attribute targets the pyir_read that precedes this call @@ -1180,8 +3119,47 @@ def _type_desc(v: Any) -> str: target_name, ) + # A staged store replacing a fold-witnessed old value invalidates the fold: + # refuse loudly (covers writes whose pre-read did not promote). + _pyir_check_staged_fold_witness(old_value, target_name) + + # Child-region escapee: the new value's SSA was produced inside an already-closed + # CF region; the in-region store carries as an iter_arg, so skip the post-store. + if reuse_accessible_ref: + skipped = _pyir_skip_trapped_child_region_writeback(mv.ref, new_value, mv=mv) + if skipped is not None: + log().info( + "[pyir_assign] '%s' new value from closed child region " + "(ref already written in-loop) -> skip redundant store", + target_name, + ) + return skipped + + _pre_store_ref = mv.ref mv.store(new_value) log().info("[pyir_assign] '%s' stored", target_name) + # Width-changing rebind: ``MutableValue.store`` re-minted the cell at the + # new MLIR type, so the name denotes a NEW place generation. Republish the + # new cell under the D1 name key when it dominates the whole function, + # else retire the entry so reads resolve the value-carried ref. + try: + if _pre_store_ref is not None and mv.ref is not None: + _same_ref = False + try: + _same_ref = bool(_pre_store_ref == mv.ref) + except Exception: + _same_ref = _pre_store_ref is mv.ref + if not _same_ref: + _wc_slot = _make_slot_key(target_name, owner, slot_name) + if _wc_slot is not None and _wc_slot in _slot_refs: + if _ref_dominates_whole_function(mv): + _slot_refs[_wc_slot] = mv.ref + _pyir_record_slot_template(_wc_slot, mv._value) + else: + del _slot_refs[_wc_slot] + _slot_templates.pop(_wc_slot, None) + except Exception: + pass # Do NOT emit pyir.load here — the loaded SSA value would be # trapped inside the current CF region. Instead, return new_value @@ -1192,7 +3170,7 @@ def _type_desc(v: Any) -> str: # Anti-aliasing: when slot context is provided ALWAYS reconstruct so # the wrapper carrying this ref belongs exclusively to this slot. # Otherwise, only reconstruct when new_value does not already carry - # this mv (the legacy anti-aliasing rule for locals). + # this mv (the anti-aliasing rule for locals). if have_slot: new_value = mv._reconstruct(new_value.ir_value()) else: @@ -1200,10 +3178,253 @@ def _type_desc(v: Any) -> str: if existing_mv is not mv: new_value = mv._reconstruct(new_value.ir_value()) _attach_mutable_ref(new_value, mv, f"pyir_assign '{target_name}'") + # The returned wrapper rewraps the just-stored raw: make it the cell's + # canonical representative so a later read proves staleness and reloads. + _pyir_adopt_stored_representative(mv, new_value) return new_value +# Attribute name carrying the ``(base, attr)`` pairs a staged loop-body +# function's ORIGINAL code assigns directly (the declared write facts). +PYIR_REGION_ATTR_WRITES_ATTR = "__pyir_region_attr_writes__" +# Attribute name carrying the ``(base, method)`` pairs a staged loop-body +# function's ORIGINAL code calls on Name-rooted receiver paths (dotted base); +# each pair completes into write facts through class facts at region entry. +PYIR_REGION_METHOD_CALLS_ATTR = "__pyir_region_method_calls__" +# Attribute name carrying the ``(func_name, arg_base)`` pairs a staged +# loop-body function's ORIGINAL code calls as a bare-Name free function with a +# Name-rooted first argument (``step(c)``); each pair completes into write +# facts through the callee's first-parameter write facts at region entry. +PYIR_REGION_FREE_CALLS_ATTR = "__pyir_region_free_calls__" + + +def pyir_tag_region_attr_writes( + attr_writes: "tuple[tuple[str, str], ...]", + method_calls: "tuple[tuple[str, str], ...]" = (), + free_calls: "tuple[tuple[str, str], ...]" = (), +) -> "Callable[[Any], Any]": + """Decorator factory tagging a staged loop-body function with the ``(base, attr)`` + pairs its ORIGINAL body assigns directly, the ``(base, method)`` pairs + it calls on Name-rooted receiver paths, and the ``(func_name, arg_base)`` + pairs it calls as free functions with a Name-rooted first argument.""" + + def _tag(func: Any) -> Any: + setattr(func, PYIR_REGION_ATTR_WRITES_ATTR, tuple(attr_writes)) + setattr(func, PYIR_REGION_METHOD_CALLS_ATTR, tuple(method_calls)) + setattr(func, PYIR_REGION_FREE_CALLS_ATTR, tuple(free_calls)) + return func + + return _tag + + +def pyir_note_attr_first_def( + target_name: str, owner: Any, slot_name: Any, value: Any +) -> None: + """Trace-only note for an attr first-def inside staged CF (the + preprocessor's first-def branch is a plain setattr, invisible otherwise). + The record is a binding-position fact, so every value kind records.""" + if owner is None or slot_name is None: + return + if not is_inside_staged_cf() or is_inside_constexpr_loop(): + return + _pyir_record_cf_attr_first_def(owner, slot_name) + + +def pyir_generation_probe(target_name: str) -> None: + """Detector-only probe for function-scope Name-rooted attribute reads: emits no + IR, evaluates no attribute, and fires only while a stamped generation exists.""" + if not ( + _SUPERSEDED_GENERATIONS + or _ALIAS_CAPTURE_ROOT_NAMES + or _SUPERSEDED_LEAF_WRAPPERS + or _PYIR_CF_ATTR_FIRST_DEFS + ): + return + if _pyir_gen_events_suppressed(): + return + root_name, sep, attr = target_name.partition(".") + if not sep or not attr or "[" in root_name: + return + base = sys._getframe(1).f_locals.get(root_name) + if base is None: + return + _pyir_generation_read_checks(target_name, None, base, attr) + _pyir_check_cf_attr_first_def_read(base, attr) + + +def _pyir_reconcile_unobserved_write( + mv: Any, + current_value: Any, + target_name: Any, + owner: Any, + slot_name: Any, +) -> "tuple[Any, Any]": + """V-3 at the read choke: the live binding is another cell's product, so a + write the chokes never saw rebound this place. Within an unchanged + region-epoch interval the write is replayed through the write choke here + (commit-on-read: the interval is straight-line, so this position is the + binding's); across a staged-region boundary the binding position is + unrecoverable -- refuse loudly.""" + if not _pyir_row_binding_unobserved_write(current_value, mv): + return mv, current_value + filename, lineno = _first_non_dsl_caller_location() + if mv._last_choke_region_epoch != _pyir_current_region_epoch(): + raise DSLUserCodeError( + DiagId.UNOBSERVED_WRITE_POSITION_UNKNOWN, + filename=filename, + lineno=lineno, + var=str(target_name), + ) + result = pyir_assign( + str(target_name), + mv._value, + current_value, + filename or "", + lineno or 0, + owner=owner, + slot_name=slot_name, + ) + if result is not None and result is not current_value: + # Python's storage adopts the choke's binding (the boundary-replay + # writeback discipline); composite spellings name no storage key. + if isinstance(owner, (dict, list)) or ( + owner is not None and isinstance(slot_name, str) + ): + _pyir_holder_store(owner, slot_name, result) + current_value = result + # The replay may have re-minted or retired the row: re-resolve the place. + if owner is not None and slot_name is not None: + mv = _get_slot_mv(owner, slot_name) + else: + mv = _get_slot_mv(None, target_name) + return mv, current_value + + +_PYIR_FABRICATION_MISS = _Sentinel("fabrication miss") + + +def _pyir_judge_fabricated_attr_read( + owner: Any, slot_name: Any, value: Any +) -> "tuple[str, str] | None": + """Fabricated-read guard (LangRef 3.12 section 3.3.2): a read FABRICATED by + a user-defined ``__getattr__`` names no storage slot, so no place fact can + carry it. A bare meta payload consumed inside staged CF would bake here + and silently mask whatever mutable state the hook derived it from -- + refuse loudly. Tolerated everywhere else: meta flow is CPython truth, + and a tracked payload (staged / watched / route-carrying / callable) + carries its own choke-governed identity through the hook. Returns the + tolerated read's ``(definer qualname, attr)`` witness (``None`` when the + read is not a tolerated fabrication).""" + if owner is None: + return None + if isinstance(slot_name, _PlaceSeg): + base = slot_name.base + if not isinstance(base, str): + return None + elif isinstance(slot_name, str): + base = slot_name + else: + return None + definer = _pyir_class_facts.getattr_fabrication_definer(type(owner)) + if definer is None: + return None + from .pyir_call_boundary import _pyir_boundary_module_is_user + + if not _pyir_boundary_module_is_user(getattr(definer, "__module__", None)): + return None # wrapper-consumer contract: DSL/stdlib fabrication protocols + if ( + _inspect_module.getattr_static(owner, base, _PYIR_FABRICATION_MISS) + is not _PYIR_FABRICATION_MISS + ): + return None # the name resolves to real storage: not a fabricated read + tracked = ( + _is_staged_value(value) + or isinstance(value, (_WatchedM, _WatchedDict, _WatchedList)) + or getattr(value, "_mutable_ref", None) is not None + or callable(value) + ) + if not tracked and is_inside_staged_cf(): + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.OWNER_FABRICATED_ATTR_IN_STAGED_CF, + filename=filename, + lineno=lineno, + obj=owner.__name__ if isinstance(owner, type) else type(owner).__name__, + attr=base, + definer=definer.__qualname__, + ) + if not tracked: + from .pyir_core import _pyir_spec_note_fabricated_bake + + # F-SPEC: an unrooted fabricated bare-meta bake is unverifiable at + # re-entry (no storage slot names a re-derivable path): fail closed. + _pyir_spec_note_fabricated_bake(owner, base) + return (definer.__qualname__, base) + + +_PYIR_MUTATING_FGET_SCANS: "dict[Any, bool]" = {} +_PYIR_FGET_STORE_OPS = frozenset( + ( + "STORE_ATTR", + "DELETE_ATTR", + "STORE_GLOBAL", + "DELETE_GLOBAL", + "STORE_SUBSCR", + "DELETE_SUBSCR", + ) +) +# A ``nonlocal`` write compiles to STORE_DEREF/DELETE_DEREF. Only a FREEvar +# target (a cell captured from an enclosing scope) is an external side effect; +# a cellvar target is a getter-local captured by a nested function and is +# benign, so the scan gates DEREF ops on ``code.co_freevars``. +_PYIR_FGET_DEREF_OPS = frozenset(("STORE_DEREF", "DELETE_DEREF")) + + +def _pyir_judge_mutating_property_read(owner: Any, slot_name: Any) -> None: + """Mutating-getter guard: a user ``@property`` whose fget's own bytecode + writes state executes once at trace time, so inside dynamic staged CF its + side effect cannot re-run per iteration -- the read result AND the mutated + state both freeze. Refuse loudly; pure getters (no store opcodes) pass, + and DSL/stdlib properties are wrapper-consumer machinery and pass.""" + if owner is None or isinstance(owner, type) or not isinstance(slot_name, str): + # A composite ``_PlaceSeg`` slot names a tuple element, never a + # descriptor attribute; only a plain-str name can resolve a property. + return + if not is_inside_staged_cf() or is_inside_constexpr_loop(): + return + descr = getattr(type(owner), slot_name, None) + if not isinstance(descr, property) or descr.fget is None: + return + fget = descr.fget + code = getattr(fget, "__code__", None) + if code is None: + return + mutates = _PYIR_MUTATING_FGET_SCANS.get(code) + if mutates is None: + from .pyir_call_boundary import _pyir_boundary_module_is_user + + mutates = _pyir_boundary_module_is_user( + getattr(fget, "__module__", None) + ) and any( + ins.opname in _PYIR_FGET_STORE_OPS + or (ins.opname in _PYIR_FGET_DEREF_OPS and ins.argval in code.co_freevars) + for ins in _dis_module.get_instructions(code) + ) + _PYIR_MUTATING_FGET_SCANS[code] = mutates + if not mutates: + return + filename, lineno = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.PROPERTY_GETTER_MUTATES_IN_STAGED_CF, + filename=filename, + lineno=lineno, + obj=type(owner).__name__, + attr=slot_name, + getter=getattr(fget, "__qualname__", slot_name), + ) + + def pyir_read( target_name: Any, current_value: Any = _PYIR_READ_SIMPLE_SENTINEL, @@ -1211,8 +3432,8 @@ def pyir_read( attach_ref: bool = True, owner: Any = None, slot_name: Any = None, - meta_only: bool = False, - **_extra_kwargs: Any, + force_promote_meta: bool = False, + row_authoritative: bool = False, ) -> Any: """Called by AST-inserted code before every read of a tracked variable. @@ -1234,7 +3455,7 @@ def pyir_read( *owner* / *slot_name* optionally identify the storage slot so ``pyir_read`` can consult the slot registry before the value-keyed - ``_mutable_ref`` cache. When either is ``None`` the legacy + ``_mutable_ref`` cache. When either is ``None`` the value-keyed value-keyed path is used unchanged. """ # Callers that omit ``current_value`` get the historical default of @@ -1242,6 +3463,50 @@ def pyir_read( # None" at the call site). if current_value is _PYIR_READ_SIMPLE_SENTINEL: current_value = None + # Refuse a wrapper minted under a previous top-level compilation before any + # dominance probe dereferences its already-finalized backing IR handle. + _pyir_guard_stale_epoch(current_value) + # F-MEMORY: an element of a memref-backed owner lives in staged memory and + # the owner's accessors emit its loads; memory, not a row, is the authority. + if ( + owner is not None + and slot_name is not None + and not isinstance(slot_name, (str, _PlaceSeg)) + and _is_memref_like(owner) + ): + return current_value + # Fabrication guard: judge a ``__getattr__``-fabricated read here (refuse a + # bare meta bake inside staged CF; tolerate + record everywhere else). + _pyir_judge_fabricated_attr_read(owner, slot_name, current_value) + # Mutating-getter guard: a user property whose fget writes state cannot + # re-run its side effect per staged iteration -- refuse inside staged CF. + _pyir_judge_mutating_property_read(owner, slot_name) + # Read-choke adoption: a plain dict read off a KNOWN holder slot adopts with + # the holder replaced; owner-less reads stay raw. + if ( + type(current_value) is dict + and owner is not None + and slot_name is not None + and _WATCHED_DICT_READ_HOOK[0] is not None + ): + current_value = _pyir_adopt_dict_value( + owner, slot_name, current_value, label=str(target_name) + ) + # List sibling: a plain list read off a KNOWN holder slot adopts with the + # holder replaced (``obj.l`` -- the attr-read hoist passes owner/slot). + if ( + type(current_value) is list + and owner is not None + and slot_name is not None + and _WATCHED_LIST_READ_HOOK[0] is not None + ): + current_value = _pyir_adopt_list_value( + owner, slot_name, current_value, label=str(target_name) + ) + # F-SPEC hop chaining: the value observed at a rooted (owner, slot) place + # roots one composition deeper, so its own leaf reads become re-derivable. + if owner is not None and slot_name is not None: + _pyir_spec_chain_value(current_value, owner, slot_name) log().info( "[pyir_read] '%s' type=%s attach_ref=%s", target_name, @@ -1249,6 +3514,19 @@ def pyir_read( attach_ref, ) + # Superseded-generation read checks: three dict truthiness tests on the + # common path; the channel checks run only while a stamp or capture exists. + if ( + _SUPERSEDED_GENERATIONS + or _SUPERSEDED_LEAF_WRAPPERS + or _ALIAS_CAPTURE_ROOT_NAMES + ): + _pyir_generation_read_checks(target_name, current_value, owner, slot_name) + # Region-conditional attr first-def read check (read-before-set): a read + # outside the first-def block's subtree is never-set on its reaching path. + if _PYIR_CF_ATTR_FIRST_DEFS and owner is not None: + _pyir_check_cf_attr_first_def_read(owner, slot_name) + # Tuple recursion: aggregate tuples are not ref-supported themselves # (``_can_create_ref(tuple) == False``), but their staged leaves often # are. ``_decompose_tuple`` creates per-element refs on the write side; @@ -1258,17 +3536,22 @@ def pyir_read( if isinstance(current_value, tuple): def elem_slot(i: int) -> Any: - return f"{slot_name}[{i}]" if slot_name is not None else None + return _place_seg_child(slot_name, i) if slot_name is not None else None - return tuple( - pyir_read( - f"{target_name}[{i}]", - elem, - attach_ref=attach_ref, - owner=owner, - slot_name=elem_slot(i), - ) - for i, elem in enumerate(current_value) + # F-TYPEID: the reload rebuilds through the declared reconstruction + # funnel so a NamedTuple keeps its class (a bare tuple() strips it). + return _rebuild_tuple_like( + current_value, + [ + pyir_read( + f"{target_name}[{i}]", + elem, + attach_ref=attach_ref, + owner=owner, + slot_name=elem_slot(i), + ) + for i, elem in enumerate(current_value) + ], ) # Slot-tracked container recursion: refresh each registered slot @@ -1280,40 +3563,49 @@ def elem_slot(i: int) -> Any: # instead of the construction-time cached SSA carried over from the # entry of the loop. # - # Gated to staged CF -- outside, the legacy code below already loads + # Gated to staged CF -- outside, the code below already loads # from any recorded ref via the slot-first lookup. Primitives are # excluded; tuples / aggregate containers handled by the branch above. # # We MUST re-attach ``_mutable_ref`` to the freshly-loaded value so a - # subsequent ``attach_ref=False`` snapshot read (e.g. ``cute.printf`` - # arg) can recover the ref via the value-keyed cache instead of + # subsequent ``attach_ref=False`` snapshot read (e.g. a downstream-DSL + # print arg) can recover the ref via the value-keyed cache instead of # falling through to ``_create_ref`` Case-D poison. if ( attach_ref and is_inside_staged_cf() and not isinstance(current_value, (int, float, bool, str, bytes)) ): - slots = _iter_slot_mvs_for_pyir_read(current_value) + slots = _iter_owner_slot_mvs(current_value) if slots: + # Refreshing a SUPERSEDED object's slots would setattr the current + # generation's values onto the retained handle; check first. + if _SUPERSEDED_GENERATIONS and not _pyir_gen_events_suppressed(): + _gen_rec = _SUPERSEDED_GENERATIONS.get(id(current_value)) + if _gen_rec is not None: + for _mv0, _v0 in _gen_rec["cells"].values(): + if _mv0._store_version > _v0: + _pyir_raise_superseded( + "read", target_name, _gen_rec["site"] + ) owner_cls = type(current_value) # Refuse silent rebind when the class overrides __setattr__: # the user's hook may have side-effects we cannot reason about, # and the alternative (calling type-specific __setattr__) would # re-trigger PyIR's instrumentation. Skip refresh in that case - # and fall through to the legacy path. - if owner_cls.__setattr__ is object.__setattr__: + # and fall through to the value-keyed path. + if _pyir_plain_storage_setattr(owner_cls): refreshed = False for stored_slot_name, slot_mv in slots: if slot_mv is None or not slot_mv._is_ref_accessible(): continue - # Skip subscript-style slot names (``'_coord[0]'``): - # those refer to tuple elements registered by - # ``_decompose_tuple`` under composite slot keys, not - # direct Python attributes. ``setattr`` on them - # would create a synthetic ``box._coord[0]`` attribute - # that then breaks ``_decompose_m2m_assign``'s - # parallel walk on the next reassignment. - if not isinstance(stored_slot_name, str) or "[" in stored_slot_name: + # Skip composite slot keys (``_PlaceSeg``): those refer + # to tuple elements registered by ``_decompose_tuple``, + # not direct Python attributes. ``setattr`` on them + # would create a synthetic attribute that then breaks + # ``_decompose_m2m_assign``'s parallel walk on the next + # reassignment. + if not isinstance(stored_slot_name, str): continue fresh = slot_mv.load() # Re-attach the slot ref to the fresh value so @@ -1325,7 +3617,7 @@ def elem_slot(i: int) -> Any: f"pyir_read '{target_name}' slot-refresh", ) try: - object.__setattr__(current_value, stored_slot_name, fresh) + _pyir_setattr_raw(current_value, stored_slot_name, fresh) refreshed = True except (AttributeError, TypeError): pass # frozen / __slots__ without target -- best-effort @@ -1350,12 +3642,14 @@ def elem_slot(i: int) -> Any: # Slot context is only authoritative when: # - caller supplied both owner and slot_name # - the owner is non-None (dict/list owners route through - # ``_slot_mvs``; ``__dict__`` owners use tier-1 / tier-2) + # ``_SLOT_REGISTRY``; ``__dict__`` owners use tier-1) # - this read will also attach _mutable_ref to the returned value # (attach_ref=False reads are snapshots -- no ref tracking needed, # and creating a ref here would introduce spurious iter_args). + # A CONTAINER owner is slot-authoritative even for snapshot reads: legs often + # share a seed value, so identity comes from the slot registry, not the value. have_slot = ( - attach_ref + (attach_ref or isinstance(owner, (dict, list))) and owner is not None and slot_name is not None and _slot_storage_available(owner) @@ -1363,14 +3657,50 @@ def elem_slot(i: int) -> Any: if not is_inside_staged_cf(): # Outside staged CF, but the value may have been modified inside - # a now-exited CF region. If a ref exists (slot-first, else on - # the value), load from it so subsequent uses see the accumulated - # result. - if have_slot: + # a now-exited CF region. If the place's ledger row exists + # (R1), load from it so subsequent uses see the accumulated + # result; a value-carried route is only place-validated evidence, + # and a computed (property-family) attribute has no row at all. + computed_slot = ( + owner is not None + and slot_name is not None + and _pyir_owner_slot_is_computed(owner, slot_name) + ) + read_place = ( + None if computed_slot else _pyir_read_place(target_name, owner, slot_name) + ) + if computed_slot: + mv = None + elif owner is not None and slot_name is not None: mv = _get_slot_mv(owner, slot_name) else: - mv = getattr(current_value, "_mutable_ref", None) + mv = _get_slot_mv(None, target_name) + if mv is not None: + # R1c (V-3): judge the live binding against the row before the + # row routes the read; an unobserved write replays or refuses. + mv, current_value = _pyir_reconcile_unobserved_write( + mv, current_value, target_name, owner, slot_name + ) + if mv is None: + route = getattr(current_value, "_mutable_ref", None) + if route is not None and ( + getattr(route, "_place", None) in (read_place, None) + if read_place is not None + else ( + _pyir_route_is_live_place_row(route) + or _pyir_route_is_current(current_value, route) + ) + ): + mv = route + if read_place is not None: + # First-wins rooting: register the adopted cell as the + # place's row so later reads resolve by place. + if owner is not None and slot_name is not None: + mv = _set_slot_mv(owner, slot_name, route) + else: + mv = _set_slot_mv(None, target_name, route) if mv is not None and mv._is_ref_accessible(): + _pyir_refuse_superseded_row_serve(mv, target_name) loaded = mv.load() log().info( "[pyir_read] '%s' outside CF but has ref → loaded", @@ -1382,15 +3712,23 @@ def elem_slot(i: int) -> Any: # D1: even outside CF, the slot may have been promoted earlier # (e.g. by a sibling for that fired _meta_promote_slot). Load # from the D1 ref so post-CF reads see the accumulated result. - d1_slot = _make_slot_key(target_name, owner, slot_name, _current_fn_id()) - if d1_slot is not None and d1_slot in _slot_refs: - return _load_as_dsl(_slot_refs[d1_slot], current_value) + d1_slot = _make_slot_key(target_name, owner, slot_name) + if ( + d1_slot is not None + and d1_slot in _slot_refs + # Skip an opaque row whose pointee no longer matches the binding: + # it belongs to an earlier generation, so the load would read stale. + and not _pyir_row_load_type_mismatch(_slot_refs[d1_slot], current_value) + ): + return _load_as_dsl(_slot_refs[d1_slot], place=d1_slot, stamp_place=True) log().info("[pyir_read] '%s' outside staged CF → passthrough", target_name) + # F-SPEC: a meta passthrough of a rooted place is a specialization fact. + _pyir_spec_record_read(read_place, current_value) return current_value # D1 (META_VALUE_TABLE_DESIGN): when inside staged CF AND the slot is # tracked (or trackable) by the meta-value table, route through D1 - # instead of the legacy paths. + # instead of the value-keyed paths. # # - If the slot is already promoted, emit ``pyir.load %ref`` and # return the matching DSL Numeric. The promotion happened earlier @@ -1398,21 +3736,23 @@ def elem_slot(i: int) -> Any: # - If the value is a Python primitive AND the slot is not yet # promoted, wrap as ``_WatchedM`` so a later mutation can rewrite # any leaf constants we bake during use. - d1_slot = ( - None - if target_name == "self.dummy" - else _make_slot_key(target_name, owner, slot_name, _current_fn_id()) - ) + d1_slot = _make_slot_key(target_name, owner, slot_name) if d1_slot is not None: existing_ref = _slot_refs.get(d1_slot) + # A2 read guard: an OPAQUE row whose pointee no longer matches the + # binding belongs to an earlier generation -- skip the load. + if existing_ref is not None and _pyir_row_load_type_mismatch( + existing_ref, current_value + ): + existing_ref = None if existing_ref is not None: - sample = ( - current_value.python_value - if isinstance(current_value, _WatchedM) - else current_value - ) - return _load_as_dsl(existing_ref, sample) - if isinstance(current_value, _WatchedM): + return _load_as_dsl(existing_ref, place=d1_slot, stamp_place=True) + # A force-promoted read must not settle for the _WatchedM wrap (a condition + # over a watched meta would fold); unwrap and fall through to Mp->S. + if force_promote_meta and isinstance(current_value, _WatchedM): + current_value = current_value.python_value + elif isinstance(current_value, _WatchedM): + _pyir_spec_record_read(current_value._slot_key, current_value) return current_value # idempotent # Strict type check (NOT isinstance): IntEnum / IntFlag subclass int, # but wrapping them in _WatchedM strips their enum identity at the @@ -1421,7 +3761,9 @@ def elem_slot(i: int) -> Any: # vs the expected ``). Only true primitives are # mutated as D1 slots; enums are configuration values and should # pass through unchanged. See PYIR_DEV_GUIDE.md Pitfall 15. - if type(current_value) in (bool, int, float): + if type(current_value) in (bool, int, float) and not force_promote_meta: + # F-SPEC: the wrapped payload is what a later bake carries. + _pyir_spec_record_read(d1_slot, current_value) return _WatchedM(current_value, d1_slot) if not _is_staged_value(current_value): @@ -1433,9 +3775,11 @@ def elem_slot(i: int) -> Any: # Only when attach_ref=True (assignment reads). Standalone # reads (attach_ref=False, e.g. self.x in a format string) # must stay as Python scalars for meta-level operations. + # ``force_promote_meta`` is the declared carried-leg fact (condition-read + # AND body-written legs must stage); it overrides the meta default. if ( attach_ref - and is_auto_m2s_enabled() + and (is_auto_m2s_enabled() or force_promote_meta) and type(current_value) in (bool, int, float) ): try: @@ -1449,7 +3793,12 @@ def elem_slot(i: int) -> Any: if mv is not None and not mv._is_ref_accessible(): mv = None if mv is None: - mv = _create_ref(promoted) + mv = _create_ref( + promoted, + birth_block=_pyir_recorded_birth_block( + target_name, owner, slot_name + ), + ) if have_slot: _set_slot_mv(owner, slot_name, mv) loaded = mv.load() @@ -1479,9 +3828,19 @@ def elem_slot(i: int) -> Any: f"pyir_read '{target_name}' Mp→S auto-promotion", ) return loaded - except Exception: - pass # No MLIR context — fall through to passthrough + except Exception as exc: + # Fail-loud wall: a swallowed promotion failure would silently + # bake the meta payload where a staged carry was declared. + raise DSLRuntimeError( + "PyIR emission self-check: Mp→S auto-promotion failed for " + f"'{target_name}' (payload {current_value!r}): {exc}" + ) from exc log().info("[pyir_read] '%s' not staged → passthrough", target_name) + # F-SPEC: a meta passthrough inside staged CF bakes wherever it is + # consumed; record it under its place when rooted. + _pyir_spec_record_read( + _pyir_read_place(target_name, owner, slot_name), current_value + ) return current_value if not _can_create_ref(current_value): @@ -1492,27 +3851,91 @@ def elem_slot(i: int) -> Any: ) return current_value - # Slot-first lookup: when owner/slot_name are supplied the slot - # registry is authoritative. Values that share a ref across slots - # (e.g. ``self.a = self.b = self.c = seed``) previously collapsed - # because ``_mutable_ref`` was attached to the shared value. - # When the slot has no entry fall back to the value-keyed - # ``_mutable_ref`` cache -- this keeps us interoperable with refs - # created by legacy code paths (``attach_ref=False`` snapshot reads, - # M->S promotions) so we do not duplicate ref storage. - if have_slot: + # R0: the read's place is derived unconditionally; attach_ref keeps its + # ref-leakage role only and never gates place validation. A computed + # (property-family) attribute names no storage place (F-SHAPE). + computed_slot = ( + owner is not None + and slot_name is not None + and _pyir_owner_slot_is_computed(owner, slot_name) + ) + read_place = ( + None if computed_slot else _pyir_read_place(target_name, owner, slot_name) + ) + # R1: a live ledger row at the place is authoritative (values that share a + # ref across slots resolve by place, never by object identity). + if computed_slot: + mv = None + elif owner is not None and slot_name is not None: mv = _get_slot_mv(owner, slot_name) - if mv is None: - legacy_mv = getattr(current_value, "_mutable_ref", None) - if legacy_mv is not None: - mv = legacy_mv - _set_slot_mv(owner, slot_name, mv) else: - mv = getattr(current_value, "_mutable_ref", None) + mv = _get_slot_mv(None, target_name) + if mv is not None and not row_authoritative: + # R1c (V-3): judge the live binding against the row before the row + # routes the read; an unobserved write replays or refuses. A + # row-authoritative caller passes a machinery re-presentation of the + # carried binding (region carry), which names no user write. + mv, current_value = _pyir_reconcile_unobserved_write( + mv, current_value, target_name, owner, slot_name + ) + foreign_stored_local_route = False + if mv is None: + route = getattr(current_value, "_mutable_ref", None) + if route is not None: + if read_place is not None: + # R2 (V-1): a value-carried cell routes a place-named read only + # when it is stamped as this place's own cell. A place-less + # cell is claimed by its first place-named read (first-wins + # rooting); a cell stamped with a DIFFERENT place is LAW-1 + # snapshot evidence and falls through to the mint path. + if getattr(route, "_place", None) in (read_place, None): + if owner is not None and slot_name is not None: + mv = _set_slot_mv(owner, slot_name, route) + else: + mv = _set_slot_mv(None, target_name, route) + else: + # Same-named local row stamped in a DIFFERENT function + # scope: candidate nonlocal-cell twin; judged at the mint + # (a snapshot passthrough below never re-roots the place). + _route_place = getattr(route, "_place", None) + if ( + owner is None + and slot_name is None + and isinstance(target_name, str) + and target_name.isidentifier() + and isinstance(_route_place, tuple) + and len(_route_place) >= 3 + and _route_place[0] == "local" + and _route_place[2] == target_name + and isinstance(read_place, tuple) + and len(read_place) >= 3 + and read_place[0] == "local" + and _route_place[1] != read_place[1] + and getattr(route, "_store_version", 0) > 0 + ): + foreign_stored_local_route = True + elif _pyir_route_is_live_place_row(route) or _pyir_route_is_current( + current_value, route + ): + # R3: a place-less read follows its route when the route IS its + # own place's live row (place authority) or while the value is + # provably the cell's current binding (V-2). + mv = route + else: + # R3 residue: the value is a retained snapshot; its own SSA is + # Python's dereference, and an unusable SSA refuses loudly. + return _pyir_resolve_snapshot(current_value, target_name) mv_was_preexisting = mv is not None log().info("[pyir_read] '%s' existing mv=%s", target_name, mv) if mv is None: + # A route-less computed (property-family) read IS its result: no + # storage place exists, so there is no cell to create or consult (R3). + if computed_slot: + log().info( + "[pyir_read] '%s' computed attribute -> passthrough", target_name + ) + return current_value # D1 bare Name-Load snapshot reads (``attach_ref=False`` AND no # owner/slot context) skip lazy ref creation -- the wrap is # only for ``_meta_uses`` recording on Python primitives. @@ -1524,10 +3947,59 @@ def elem_slot(i: int) -> Any: target_name, ) return current_value - mv = _create_ref(current_value) + # An assignment pre-read promoting a fold-witnessed value to a cell is + # the write transition: refuse before the stale fold carries a store. + if attach_ref: + _pyir_check_staged_fold_witness(current_value, target_name) + # Backstop: a ``nonlocal`` name resolves to its cell's home row + # (see ``_cell_home_binding``), so a foreign stored row here means + # the home fact is absent -- a re-rooting mint would abandon the + # foreign scope's stores at the join. Refuse instead of twinning. + if ( + foreign_stored_local_route + and is_inside_staged_cf() + and _PYIR_SCOPE_STACK + and target_name + in _PYIR_SCOPE_NONLOCAL_NAMES.get(_PYIR_SCOPE_STACK[-1].scope_id, ()) + ): + raise DSLUserCodeError( + DiagId.SCOPE_NONLOCAL_WRITE_IN_STAGED_CF, + name=str(target_name), + ) + # R6: a row-authoritative carry mint firing inside the promoted loop + # op's own block emits its cell and copy-in at the declared pre-region + # position (write facts are position + value; never the ambient IP). + _pre_ip = _pyir_promoted_region_pre_ip() if row_authoritative else None + _birth = _pyir_recorded_birth_block(target_name, owner, slot_name) + # Fail-closed sibling-region serve guard: a place-named container/attr + # read minting from an SSA trapped in a closed region that does not + # dominate this IP serves a value Python never bound on this path + # (e.g. a protocol-carried holder rebuilt at region close presents the + # producer arm's trace-time value to the sibling arm). A recorded + # birth block keeps the region-born cell mint (seed placed at birth). + if ( + owner is not None + and slot_name is not None + and _birth is None + and _raw_backing_ir_value(current_value) is not None + and not _value_dominates_current_ip(current_value) + ): + _sg_file, _sg_line = _first_non_dsl_caller_location() + raise DSLUserCodeError( + DiagId.SCOPE_READ_NEVER_SET, + filename=_sg_file, + lineno=_sg_line, + detail=" (its only value was produced inside a sibling branch" + " that cannot reach this read)", + ) + if _pre_ip is not None: + with _pre_ip: + mv = _create_ref(current_value, birth_block=_birth) + else: + mv = _create_ref(current_value, birth_block=_birth) log().info("[pyir_read] '%s' ref created via _create_ref", target_name) if have_slot: - _set_slot_mv(owner, slot_name, mv) + mv = _set_slot_mv(owner, slot_name, mv) # Reconstruct BEFORE attach when slot context is present so the # fresh wrapper carries this ref instead of the caller's # (potentially shared) value. Without this, a later write to a @@ -1537,9 +4009,81 @@ def elem_slot(i: int) -> Any: # value-keyed ``_mutable_ref`` cache is the only lookup path # available to subsequent reads. current_value = mv._reconstruct(current_value.ir_value()) + elif ( + owner is not None + and slot_name is not None + and is_inside_staged_cf() + and _slot_storage_available(owner) + ): + # PRODUCER PUBLISH. + mv = _set_slot_mv(owner, slot_name, mv) + # LAW-2: a published place row's representative is a fresh + # wrapper -- the presented object can be (or become) a live + # binding of ANOTHER place (a frame local, a kernel parameter), + # and attaching this row to it would steal that binding's route. + # A literal-backed value mints a fresh position-independent + # constant per ``ir_value()``; an adoptable slot mints it at the + # declared pre-region position so the constant lands outside the + # staged region. Every other value materializes position- + # dependently (e.g. as a cell load) and keeps the ambient mint. + # Composite spellings name no storage key. + _adopt_slot = isinstance(slot_name, str) + _pub_pre_ip = ( + _pyir_promoted_region_pre_ip() + if _adopt_slot and _is_literal_backed(current_value) + else None + ) + if _pub_pre_ip is not None: + with _pub_pre_ip: + _pub_raw = current_value.ir_value() + else: + _pub_raw = current_value.ir_value() + current_value = mv._reconstruct(_pub_raw) + # Python's storage adopts the representative (the boundary-replay + # writeback discipline: the slot's binding stays routed to its + # row, so later slot reads and boundary-observed writes resolve + # it) only when every later serve of the adopted binding is + # total: the raw itself dominates the region exterior (a direct + # serve is valid anywhere), or the row's cell does (every choke + # read re-loads the live cell and the boundary replay carries + # it -- a plain-method advance inside the region then reads and + # writes the carried cell, not a baked snapshot). A raw AND + # cell both interior to a staged region stay out of storage: + # the row alone carries the binding. + _row_cell = getattr(mv, "_ref", None) + if _adopt_slot and ( + _pyir_raw_dominates_region_exterior(_pub_raw) + or ( + _row_cell is not None + and _pyir_raw_dominates_region_exterior(_row_cell) + ) + ): + # The stored wrapper IS the row's own representative: the cell + # was minted from this very value at this choke, so a later + # region-close sweep finding it unchanged in storage has no + # un-instrumented advance to store back. A real plain-method + # advance rebinds the slot to a NEW object without this mark. + try: + _pyir_setattr_raw( + current_value, "_pyir_publish_representative_of", mv + ) + except (AttributeError, TypeError): + pass # unmarkable wrapper: the sweep judges it by position + _pyir_holder_store(owner, slot_name, current_value) + elif owner is None and read_place is not None: + # R2 publish: the minted snapshot cell becomes the local place's + # row, so later same-name reads resolve by place (R1), not by value. + mv = _set_slot_mv(None, target_name, mv) + # LAW-2: a wrapper routed to a FOREIGN cell is a shared caller-scope + # binding; this scope's representative is minted on a fresh wrapper. + _prior_route = getattr(current_value, "_mutable_ref", None) + if _prior_route is not None and _prior_route is not mv: + current_value = mv._reconstruct(current_value.ir_value()) _attach_mutable_ref(current_value, mv, f"pyir_read '{target_name}' new-ref") if mv._is_ref_accessible(): + if mv_was_preexisting: + _pyir_refuse_superseded_row_serve(mv, target_name) loaded = mv.load() log().info( "[pyir_read] '%s' loaded → %s", @@ -1558,6 +4102,10 @@ def elem_slot(i: int) -> Any: _attach_mutable_ref(loaded, mv, f"pyir_read '{target_name}' loaded") return loaded + # A SNAPSHOT read of a region-trapped container leg must not re-create the + # cell (the write side owns that); the trapped-raw recovery repairs reads. + if not attach_ref and isinstance(owner, (dict, list)): + return current_value # Ref exists but is inaccessible (defined in a sibling/exited CF region, # e.g. inside a previous meta-loop iteration's scf.if body). # current_value's SSA may not dominate the current insertion point. @@ -1565,9 +4113,22 @@ def elem_slot(i: int) -> Any: # _create_ref handles placement: literal-backed → function entry (Case A), # SSA-backed → after defining op or current IP (Cases B-D). try: - mv = _create_ref(current_value) - if have_slot: + if attach_ref: + _pyir_check_staged_fold_witness(current_value, target_name) + # R6: same declared pre-region position for the re-created cell of a + # row-authoritative carry (see the mint above). + _pre_ip = _pyir_promoted_region_pre_ip() if row_authoritative else None + _birth = _pyir_recorded_birth_block(target_name, owner, slot_name) + if _pre_ip is not None: + with _pre_ip: + mv = _create_ref(current_value, birth_block=_birth) + else: + mv = _create_ref(current_value, birth_block=_birth) + if have_slot and not computed_slot: _set_slot_mv(owner, slot_name, mv) + elif owner is None and read_place is not None: + # F-PLACE: the re-created cell becomes the local place's row. + mv = _set_slot_mv(None, target_name, mv) loaded = mv.load() log().info( "[pyir_read] '%s' ref inaccessible → re-created + loaded", @@ -1577,12 +4138,178 @@ def elem_slot(i: int) -> Any: _attach_mutable_ref(loaded, mv, f"pyir_read '{target_name}' re-created") return loaded except DSLUserCodeError: - # Don't swallow user-facing diagnostics (e.g. poison-read catcher). - raise + # Don't swallow user-facing diagnostics (e.g. poison-read catcher). + raise + except Exception: + # If re-creation fails (e.g. non-ref-compatible type), passthrough. + log().info("[pyir_read] '%s' ref not accessible → passthrough", target_name) + return current_value + + +def pyir_bind_param(target_name: str, value: Any) -> Any: + """Scope entry: a parameter binding IS this scope's first-def (Python + semantics), so a binding that arrives routed to ANOTHER place's live row + is re-homed now -- the scope's own row is minted on a fresh wrapper + seeded from the binding (LAW-2), and every later read, write, and + boundary load in this scope resolves the scope's own place row instead + of the caller's shared cell. Route-less or own-row bindings keep the + lazy row mint. A container param binds each leaf place the same way.""" + if isinstance(value, tuple): + return _pyir_bind_param_leaves(target_name, value) + return _pyir_bind_param_scalar(target_name, value) + + +def _pyir_bind_param_scalar(target_name: str, value: Any) -> Any: + """The scalar arm of :func:`pyir_bind_param` (one name, one place).""" + route = getattr(value, "_mutable_ref", None) + if route is None: + # A route-less staged binding shares the caller's wrapper OBJECT; the + # callee owns its binding (LAW-2), so hand it a fresh wrapper -- else a + # later callee row mints on the shared wrapper and the rebind leaks out. + if _is_staged_value(value) and _can_create_ref(value): + return _fresh_wrapper(value) + return value + place = _pyir_read_place(target_name, None, None) + if place is None: + return value + route_place = getattr(route, "_place", None) + if route_place is None or route_place == place: + return value + if not _can_create_ref(value) or not hasattr(value, "ir_value"): + return value + # The argument's value AT THE BIND POSITION is a read of the caller's live + # cell (the wrapper's cached SSA may predate in-region advances of that + # cell), so seed from a route load emitted here. + incoming = route.load() + mv = _create_ref(incoming) + mv = _set_slot_mv(None, target_name, mv) + # The binding IS the scope's first-def: commit the incoming value into the + # scope's own row AT THE BIND POSITION (inside the current region for an + # inlined callee, so a loop re-binds per iteration). Row-keyed reads then + # resolve this write instead of the row's placeholder init. + mv.store(incoming) + fresh = mv._reconstruct(incoming.ir_value()) + _attach_mutable_ref(fresh, mv, f"pyir_bind_param '{target_name}'") + return fresh + + +def _pyir_bind_param_leaves(target_name: str, value: tuple) -> tuple: + """A tuple/NamedTuple parameter binding is a first-def of EVERY leaf place + it creates (``name[i]``, the write-walk's leaf naming): a routed leaf is + re-homed onto this scope's own leaf row exactly like the scalar case; a + meta or route-less staged leaf records its bind position so the lazy cell + mint seeds HERE and zero-trip/post-loop reads stay defined.""" + rebuilt: list = [] + changed = False + for i, leaf in enumerate(value): + leaf_name = f"{target_name}[{i}]" + if isinstance(leaf, tuple): + new_leaf = _pyir_bind_param_leaves(leaf_name, leaf) + else: + pyir_seed_param_bindings(leaf_name) + _pyir_record_bind_position_first_def(leaf_name, leaf) + new_leaf = _pyir_bind_param_scalar(leaf_name, leaf) + changed = changed or new_leaf is not leaf + rebuilt.append(new_leaf) + if not changed: + return value + return _rebuild_tuple_like(value, rebuilt) + + +def _pyir_record_bind_position_first_def(leaf_name: str, leaf: Any) -> None: + """Record the bind position as *leaf_name*'s first-def fact (F-BIRTHPOS): + a later lazy mint / M->S promotion of the leaf place seeds at this block + instead of a hoisted placeholder. Recording only -- no wrap, no ref.""" + try: + slot = _make_slot_key(leaf_name, None, None) + except Exception: + return + if slot is None or slot in _slot_refs: + return + try: + block = ir.InsertionPoint.current.block + except Exception: + return + if _is_staged_value(leaf) and _can_carry_leaf_ref(leaf): + _slot_first_def_block[slot] = block + return + if type(leaf) in (bool, int, float): + inside_cf = is_inside_staged_cf() + _slot_first_def_inside_cf[slot] = inside_cf + if inside_cf: + _slot_first_def_depth[slot] = current_staged_cf_depth() + _slot_first_def_block[slot] = block + + +def _pyir_obs_read( + target_name: str, base: Any, attr: str, mod_name: "str | None" = None +) -> Any: + """Record-only read choke for attribute reads the staging chokes never + route (function-scope symbol-rooted reads, test-position reads, and + module/class-spelled reads). Evaluates the read exactly in place, adopts + dict/list values off the known holder slot (the object-level chokes then + own their entry lifecycles), and records the baked payload under the + read's place (F-SPEC). Never alters staging semantics.""" + value = getattr(base, attr) + try: + if base is None or _is_staged_value(base) or isinstance(base, _WatchedM): + return value + if ( + type(value) is dict + and _WATCHED_DICT_READ_HOOK[0] is not None + and not isinstance(base, (types.ModuleType, type)) + ): + value = _pyir_adopt_dict_value(base, attr, value, label=str(target_name)) + elif ( + type(value) is list + and _WATCHED_LIST_READ_HOOK[0] is not None + and not isinstance(base, (types.ModuleType, type)) + ): + value = _pyir_adopt_list_value(base, attr, value, label=str(target_name)) + # A test-position read of a meta primitive off a PLAIN compound folds + # the payload into a trace-time decision with no retargetable SSA: + # witness it like a structural consumption, so a later carry/promote + # of the place with a CHANGED value refuses instead of keeping the + # folded branch on every iteration. + if ( + type(value) in (bool, int, float) + and not isinstance(base, (types.ModuleType, type)) + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + and _innermost_enclosing_loop_op_at_ip() is not None + ): + _obs_sk = _make_slot_key(None, base, attr) + if ( + _obs_sk is not None + and _obs_sk not in _PYIR_STRUCTURAL_META_CONSUMPTIONS + ): + _obs_file, _obs_line = _first_non_dsl_caller_location() + _PYIR_STRUCTURAL_META_CONSUMPTIONS[_obs_sk] = ( + str(value), + _obs_file or "", + _obs_line or 0, + ) + _pyir_spec_observe_attr_read(target_name, base, attr, value, mod_name) + # The module/class-spelled funnel judges fabricated reads too + # (the metaclass ``__getattr__`` leg arrives here); after observation + # so a symbol-rooted base resolves its root before the bake note. + _pyir_judge_fabricated_attr_read(base, attr, value) + except DSLUserCodeError: + raise # curated refusals are the loud floor, never swallowed except Exception: - # If re-creation fails (e.g. non-ref-compatible type), passthrough. - log().info("[pyir_read] '%s' ref not accessible → passthrough", target_name) - return current_value + pass # observation must never break the read (recording fails closed) + return value + + +def _pyir_obs_global_read(name: str, value: Any, mod_name: "str | None" = None) -> Any: + """Record-only read choke for a bare global-name read (no staging choke + exists for module-level bindings): records scalar bakes under the module + root and roots object bindings for their downstream legs (F-SPEC).""" + try: + _pyir_spec_observe_global_read(name, value, mod_name) + except Exception: + pass # observation must never break the read (recording fails closed) + return value def _pyir_post_subscript_read( @@ -1593,45 +4320,93 @@ def _pyir_post_subscript_read( """Read ``container[key]``, emitting ``pyir.load`` for tracked values. Called by AST-inserted code for every ``container[key]`` in Load - context inside ``@cute.jit`` bodies. + context inside jit-decorated bodies. For dicts/lists: evaluates ``container[key]`` and passes through ``pyir_read`` with ``attach_ref=False`` to emit ``pyir.load`` when the value carries a ``_mutable_ref``. Returns the loaded value (fresh SSA) or the original value if no ref exists. + For a DICT container the owner/key pair is forwarded into ``pyir_read`` + so the read names the same ``("subscript", owner_token, key)`` D1 slot the + write choke names (the write side has always passed them) -- place + NAMING only: the read stays ``attach_ref=False`` (a standalone snapshot + read), it just stops keying the entry by its spelling string. Without + the owner, a read spelled ``d['x']`` and a read spelled ``obj.d['x']`` + of the SAME entry land on two unrelated name slots and a write-side + promotion can never find the reads' baked constants. + For non-dicts (GPU arrays, tensors, etc.): evaluates and returns ``container[key]`` directly (no overhead beyond the isinstance check). """ + # A meta KEY indexing any container is a structural consumption of the + # key's place: C-level int-subclass fast paths (tuple/list subscript via + # ``PyNumber_Index``) skip ``__index__``, so the choke must witness it -- + # records the bake and refuses an arm-local escape. + if isinstance(key, _WatchedM): + key._record_structural_consumption() + # Dict SUBCLASS with a MISSING key inside dynamic staged CF: subclass magic + # would create the key in-region (key-set rule); guarded BEFORE the access. + if ( + isinstance(container, dict) + and type(container) is not dict + and not isinstance(container, _WatchedDict) + and not dict.__contains__(container, key) + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + ): + raise DSLUserCodeError( + DiagId.CONTAINER_DICT_KEY_SET_MUTATED, + var=str(target_name), + detail=f"key {key!r} created via {type(container).__name__}", + ) val = container[key] + if isinstance(container, (_WatchedDict, _WatchedList)): + # The watched ``__getitem__`` already routed this read (single-fire + # contract); a second pyir_read would duplicate the load/wrap. + return val if not isinstance(container, (dict, list)): + # A tuple slot cannot re-point (no adoption): the item hop still + # declares the value's root chain so deeper leaf reads re-derive. + if ( + type(container) is tuple + and isinstance(key, int) + and not isinstance(key, bool) + ): + idx = key + len(container) if key < 0 else key + _pyir_spec_chain_value(val, container, steps=(("item", idx),)) return val + if isinstance(container, dict): + if type(val) is dict: + # A nested plain dict surfacing at the read choke: adopt it as a + # child place (the parent container IS the holder slot in hand). + return _pyir_adopt_dict_value(container, key, val, label=target_name) + if type(val) is list: + # A nested plain list surfacing at the read choke: same child-place + # adoption, integer-leg domain (dict-of-lists). + return _pyir_adopt_list_value(container, key, val, label=target_name) + # A plain OBJECT surfacing through a never-adopted dict's leg: its + # attr legs are still places one composition deeper. + _pyir_register_container_held_object(val) + return pyir_read( + target_name, val, attach_ref=False, owner=container, slot_name=key + ) + if type(val) is dict: + # list-of-dicts: the list element slot is the holder in hand. + return _pyir_adopt_dict_value(container, key, val, label=target_name) + if type(val) is list: + return _pyir_adopt_list_value(container, key, val, label=target_name) + if isinstance(key, int) and not isinstance(key, bool): + # Owner-keyed read of an integer leg so read and write chokes name the + # SAME subscript slot; the read itself stays a standalone snapshot. + _pyir_register_container_held_object(val) + idx = key + len(container) if key < 0 else key + return pyir_read( + target_name, val, attach_ref=False, owner=container, slot_name=idx + ) return pyir_read(target_name, val, attach_ref=False) -def _subscript_container_is_dsl_managed(container: object) -> bool: - """Whether *container*'s type builds its subscript writes as IR directly. - - DSL-managed containers (cutlass.Array, tensors, shared memory) lower - ``container[k] = v`` to real IR store ops, so they stay on the untracked - native path. Plain Python mutable containers (bytearray / array.array / - collections.deque / an author class with __setitem__) fold the write at - trace time -- inside staged CF a silent freeze -- so they are rejected. - is_dsl_internal_code alone misclassifies stdlib types (bytearray in - builtins, array in array, deque in collections are all in - sys.stdlib_module_names), so exclude stdlib explicitly. - """ - tp = type(container) - module_name = getattr(tp, "__module__", "") or "" - top = module_name.split(".", 1)[0] - if top and top in getattr(sys, "stdlib_module_names", frozenset()): - return False - # DSL-managed containers live in the ``cutlass`` DSL package; author code - # (module ``__main__`` / a user file) and stdlib are not managed. Self - # contained approximation of ``is_dsl_internal_code`` (which master lacks). - return top == "cutlass" - - def _pyir_pre_subscript_assign( target_name: str, container: object, @@ -1640,7 +4415,7 @@ def _pyir_pre_subscript_assign( """Read the old value from a subscript target for PyIR instrumentation. Called by AST-inserted code before every ``container[key] = expr`` and - ``container[key] += expr`` inside ``@cute.jit`` bodies. + ``container[key] += expr`` inside jit-decorated bodies. Returns ``_PYIR_SKIP`` if the container is not a tracked Python container (e.g. GPU array, tensor, shared memory). @@ -1650,175 +4425,950 @@ def _pyir_pre_subscript_assign( remain meta values so trace-time counters can still drive ``const_expr`` dispatch. - Lists are only tracked when the current ``@cute.jit`` body opened + Lists are only tracked when the current jit-decorated body opened staged CF. This avoids instrumenting trace-time list fills inside - nested ``@cute.jit`` helpers while still lifting list slot writes that + nested jit-decorated helpers while still lifting list slot writes that would otherwise leak SSA across ``scf.if``/``scf.for`` regions. If the key does not yet exist in the container (first-time definition), returns ``None`` so that ``pyir_assign`` treats it as a first def. + + A :class:`_WatchedDict` / :class:`_WatchedList` container returns + ``_PYIR_SKIP``: the object-level choke owns the whole write lifecycle. + The AST expansion stores through the container up to three times + (old-value writeback, the original statement's store, the + processed-value store); deferring to the watched ``__setitem__`` -- + which fires exactly once per plain store, on the original statement -- + keeps the access single-fired. + + A memref-backed container also returns ``_PYIR_SKIP``, before any + element access: element places of memref-backed owners are + memory-authoritative (the same declared fact that gates the + ``pyir_assign``/``pyir_read`` pass-throughs), so no row exists and + no instrumentation access may be emitted. """ - if isinstance(container, list): - if not is_inside_locally_staged_cf(): - return _PYIR_SKIP - elif not isinstance(container, dict): - # A plain Python mutable container (bytearray / array.array / deque, or - # an author class with __setitem__) folds the subscript write at trace - # time; inside auto-M2S staged CF that silently freezes it. A - # DSL-managed container lowers to real IR and stays on the native path. - # Constexpr scopes realize the write deterministically -> exempt. - if ( - is_auto_m2s_enabled() - and is_inside_staged_cf() - and not is_inside_constexpr_loop() - and not _subscript_container_is_dsl_managed(container) - ): + # F-MEMORY: an element of a memref-backed owner lives in staged memory + # and the owner's accessors emit its loads/stores; memory, not a row, + # is the authority -- skip before the first element access. + if not isinstance(key, str) and _is_memref_like(container): + return _PYIR_SKIP + if isinstance(container, (_WatchedDict, _WatchedList)): + return _PYIR_SKIP + # Untracked Python-mutable sequences and dict/list SUBCLASSES never become + # watched place owners: element writes inside dynamic staged CF are refused. + _untracked = isinstance( + container, (bytearray, _array_module.array, _collections_module.deque) + ) or (isinstance(container, list) and type(container) is not list) + if _untracked: + # Human name computed from the type (module-qualified for stdlib + # containers), never spelled as a literal. + _kind_ty = type(container) + _untracked_kind = _kind_ty.__name__ + if _kind_ty.__module__ not in ("builtins", None): + _untracked_kind = _kind_ty.__module__ + "." + _untracked_kind + if is_inside_staged_cf() and not is_inside_constexpr_loop(): raise DSLUserCodeError( DiagId.CONTAINER_SUBSCRIPT_WRITE_UNTRACKED, - py_type=type(container).__name__, + var=str(target_name), + kind=_untracked_kind, ) return _PYIR_SKIP + # Dict SUBCLASS with a MISSING key at the write: subclass magic materialises + # the entry (a key-set change, loud in dynamic CF); existing keys stay tracked. if ( - type(container) is not dict - and hasattr(type(container), "__missing__") + isinstance(container, dict) + and type(container) is not dict + and not dict.__contains__(container, key) and is_inside_staged_cf() and not is_inside_constexpr_loop() ): - # A dict SUBCLASS with __missing__ (Counter, defaultdict) materialises - # keys implicitly on access; that key-set change cannot be reproduced - # in staged IR. Plain dict / OrderedDict (no __missing__) stay tracked. - raise DSLUserCodeError(DiagId.CONTAINER_DICT_KEY_SET_MUTATED, var=target_name) + raise DSLUserCodeError( + DiagId.CONTAINER_DICT_KEY_SET_MUTATED, + var=str(target_name), + detail=f"key {key!r} created via {type(container).__name__}", + ) + if isinstance(container, list): + if not is_inside_locally_staged_cf(): + return _PYIR_SKIP + elif not isinstance(container, dict): + # Owner classification (funnel totality): a staged owner or a declared + # DynamicExpression compound emits/reconstructs its own element writes. + if _is_staged_value(container) or _implements_dynamic_expression(container): + return _PYIR_SKIP + # Anything else stores through an opaque ``__setitem__`` into state no + # place row names. Pre-existing reachable leaves are carried by the + # region walks; a container key CREATED by the write has no + # pre-region cell, so it can never follow the region's runtime + # predicate -- arm a region-close key-set audit of the owner. + if is_inside_staged_cf() and not is_inside_constexpr_loop(): + _pyir_arm_opaque_owner_audit(str(target_name), container) + return _PYIR_SKIP try: old = container[key] # type: ignore[call-overload] except (KeyError, IndexError): return None if isinstance(container, dict) and isinstance(old, (bool, int, float)): return old - # Pass owner/slot_name so the slot registry distinguishes subscript - # entries that happen to hold the same Python value. Without this, - # three dict keys holding the same Int32 seed would collapse to a - # single ``pyir.ref`` keyed off the shared value. + # Pass owner/slot_name so the slot registry distinguishes subscript entries that happen to hold the same + # Python value; otherwise dict keys sharing one value would collapse onto a single ``pyir.ref``. return pyir_read(target_name, old, owner=container, slot_name=key) -def with_ctxmgr_check(cm: _CM) -> _CM: - """Trace-time guard for ``with`` context managers inside staged CF. +def _attr_owner_plain_storage_setattr(owner: object) -> bool: + """Plain-storage judgment for an ATTRIBUTE assignment *owner*. + + ``_pyir_plain_storage_setattr`` inspects ``type(owner).__setattr__``. For + a CLASS owner (``cls.attr = ...``) that slot is the metaclass's, and + ``type.__setattr__`` is not ``object.__setattr__`` even though an + unadorned class stores attributes through plain type-dict storage -- so + normal class-attribute writes would be misclassified as non-plain and a + same-object write would wrongly veto its generated write-back. Treat an + unoverridden metaclass ``__setattr__`` as plain; a metaclass that DOES + override it keeps the skip protocol like any other custom hook. + """ + if isinstance(owner, type): + return type(owner).__setattr__ is type.__setattr__ + return _pyir_plain_storage_setattr(type(owner)) - A ``with`` whose ``__enter__``/``__exit__`` are plain Python user functions - runs those dunders ONCE while the kernel is traced. Inside a staged - for/while/if the block would (in Python) run them on every pass / only on - the taken path, so their effects on hidden state are frozen at the first - trace -- a silent miscompile. Reject that; pass through decorated dunders - (jit-traced per pass) and DSL-internal / stdlib context managers - (e.g. ``contextlib.suppress``). - A context-manager class carrying the ``__dsl_trace_time_ctxmgr__ = True`` - opt-in marker asserts its dunders are trace-time-only bookkeeping and is - waved through: for such a manager, running the dunders once at trace time - is the correct semantics. +def _pyir_pre_attr_assign( + target_name: str, + current_value: object, + *, + owner: object = None, + slot_name: object = None, +) -> object: + """Pre-store refresh for an ATTRIBUTE assignment target -- the attribute + sibling of ``_pyir_pre_subscript_assign``'s ``_PYIR_SKIP`` protocol. - Returns *cm* unchanged so ``with with_ctxmgr_check(EXPR) as x`` behaves - exactly like ``with EXPR as x`` in every accepted case. + The AST expansion stores this function's result straight back through + *owner*'s REAL ``__setattr__`` (``obj.attr = ``), so the result + must be a value that store can accept. ``pyir_read`` runs first with all + its tracking side effects intact; the write-back is then vetoed + (``_PYIR_SKIP``) exactly when no fresh value was produced AND the owner's + ``__setattr__`` is not plain object storage. For such an owner the + round-trip ``obj.attr = obj.attr`` is not a safe no-op: ``getattr`` may + expose a handle the custom hook cannot accept as a stored value (a + downstream-DSL struct field reads back as its scalar-slot pointer + handle; assigning that handle back raises in the DSL's numeric + conversion). When ``pyir_read`` DID + produce a fresh value the write-back proceeds unchanged -- a fresh load + is a genuine value the owner's hook must accept anyway. """ - try: - from .multi_stage_manager import is_inside_staged_cf - except Exception: # noqa: BLE001 -- staging machinery absent: nothing to guard - return cm - # A top-level ``with`` runs its dunders exactly once at trace time, which - # matches Python; only staged re-execution / conditional execution freezes - # them. + fresh = pyir_read(target_name, current_value, owner=owner, slot_name=slot_name) + if ( + fresh is current_value + and owner is not None + and not _attr_owner_plain_storage_setattr(owner) + ): + return _PYIR_SKIP + return fresh + + +def _pyir_post_attr_assign( + target_name: str, + old_value: object, + new_value: object, + filename: "str | None", + lineno: "int | None", + *, + owner: object = None, + slot_name: object = None, +) -> object: + """Post-store bookkeeping for an ATTRIBUTE assignment target. + + ``pyir_assign`` runs first and unconditionally, mirroring + ``_pyir_pre_attr_assign``'s ``pyir_read``; the sentinel vetoes only the + generated write-back ``obj.attr = ``. + + *old_value* / *new_value* are the attribute read back around the store. + Identical reads on a non-plain owner OVER-APPROXIMATE the case the veto + is for -- a downstream-DSL struct field reads back as its scalar-slot + pointer handle, which the real ``__setattr__`` rejects in its numeric + conversion. They also catch a verbatim ``object.__setattr__`` forwarder + , where the vetoed write-back would only + re-store what ``getattr`` just returned. + """ + result = pyir_assign( + target_name, + old_value, + new_value, + filename, + lineno, + owner=owner, + slot_name=slot_name, + ) + if ( + new_value is old_value + and owner is not None + and not _attr_owner_plain_storage_setattr(owner) + ): + log().info( + "[pyir_assign] '%s' handle-shaped attr on %s -- write-back skipped", + target_name, + type(owner).__name__, + ) + return _PYIR_SKIP + return result + + +def _promote_carried_attr_legs( + root_name: str, + root_obj: object, + attr_paths: "tuple[tuple[str, ...], ...]", + *, + force_promote_meta: bool, +) -> None: + """Route each declared carried attr leg through ``pyir_read`` with its + ``(owner, slot)`` identity and rebind it (default ``__setattr__`` only).""" + for path in attr_paths: + if not path: + continue + owner: Any = root_obj + skip = False + for hop in path[:-1]: + try: + owner = getattr(owner, hop) + except AttributeError: + skip = True + break + if owner is None: + skip = True + break + if skip or owner is None: + continue + leaf = path[-1] + if not isinstance(leaf, str): + continue + owner_cls = type(owner) + if not _pyir_plain_storage_setattr(owner_cls): + log().info( + "[pyir while] carried attr leg %s.%s skipped: %s overrides __setattr__", + root_name, + ".".join(path), + owner_cls.__name__, + ) + continue + try: + value = getattr(owner, leaf) + except AttributeError: + continue + dotted = f"{root_name}.{'.'.join(path)}" + fresh = pyir_read( + dotted, + value, + owner=owner, + slot_name=leaf, + force_promote_meta=force_promote_meta, + ) + if fresh is not value: + try: + _pyir_setattr_raw(owner, leaf, fresh) + except (AttributeError, TypeError): + pass # frozen / __slots__ without target -- best-effort + + +def pyir_promote_while_carried_arg( + target_name: str, + current_value: object, + is_condition_read: bool = True, + condition_attr_paths: "tuple[tuple[str, ...], ...]" = (), +) -> object: + """Materialise a ``scf.while`` write_arg's ref before the condition runs (an + unpromoted meta write_arg would fold to ``scf.condition(%true)``).""" if not is_inside_staged_cf(): - return cm - from .common import is_dsl_internal_code # inline: break the import cycle - from .ast_helpers import _is_dsl_traced_callable - - tp = type(cm) - - # Explicit opt-in: a manager whose dunders do only trace-time bookkeeping - # is correct to run once at trace time. - if getattr(tp, "__dsl_trace_time_ctxmgr__", False): - return cm - - # A generator-based ``@contextmanager`` hides its trace-time-frozen work in - # the USER generator body, while its ``__enter__``/``__exit__`` live in - # ``contextlib`` (stdlib) -- so the type-dunder check below would wrongly - # wave it through. Classify by the generator's own origin (``cm.gen``): a - # user-code generator run inside staged CF freezes its ``yield``-straddling - # side effects at the first trace and must be rejected; an internal / stdlib - # generator manager is allowed. - gen = getattr(cm, "gen", None) - gen_code = getattr(gen, "gi_code", None) - if gen_code is not None: - gen_frame = getattr(gen, "gi_frame", None) - gen_module = "" - if gen_frame is not None: - gen_module = gen_frame.f_globals.get("__name__", "") or "" - if not is_dsl_internal_code(getattr(gen_code, "co_filename", None), gen_module): + return current_value + # Promote declared carried attr legs to their place cells BEFORE the + # condition evaluates, so it reads the carried cell. + if ( + condition_attr_paths + and current_value is not None + and not isinstance(current_value, (bool, int, float, str, bytes)) + and _has_instance_storage(current_value) + ): + _promote_carried_attr_legs( + target_name, current_value, condition_attr_paths, force_promote_meta=True + ) + if ( + not isinstance(current_value, (bool, int, float)) + and current_value is not None + and _has_instance_storage(current_value) + and _implements_dynamic_expression(current_value) + and not _can_create_ref(current_value) + and not _is_scalar_ssa_carryable(current_value) + ): + _promote_compound_leaf_refs(current_value, target_name) + return current_value + if isinstance(current_value, (bool, int, float)): + if not is_condition_read: + # A meta-primitive the condition does not read is rederived per + # iteration, not loop-carried; promoting would demote a constexpr. + return current_value + return _promote_loop_carried_meta(target_name, current_value) + if ir is not None and isinstance(current_value, ir.Value): + _pyir_refuse_stale_raw_carry(target_name, current_value, "`while` condition") + if is_condition_read and isinstance( + current_value.type, (ir.IntegerType, ir.FloatType) + ): + # A raw scalar SSA local the condition reads is a loop-carried + # place: mint its cell (the body's write choke stores raw SSA + # rebinds through the same row) and serve the condition from it, + # so ``scf.condition`` reads the carried arg, not the stale + # trace-time SSA. + return _pyir_promote_raw_scalar_carry(target_name, current_value) + # The carried binding is the carry engine's re-presentation of the + # place, not a user rebind: the row stays authoritative for this read. + return pyir_read(target_name, current_value, row_authoritative=True) + + +def _pyir_refuse_stale_raw_carry( + target_name: str, current_value: "ir.Value", role: str +) -> None: + """Refuse a raw SSA binding from an already-finalized compilation: it + cannot be carried (its backing IR is gone). Runs before any dereference + (only ``.context`` is a safe probe on a possibly-dangling value).""" + try: + _raw_ctx = id(current_value.context) + except Exception: + _raw_ctx = None + if _raw_ctx != id(ir.Context.current): + raise DSLRuntimeError( + f"the {role} variable '{target_name}' holds an IR " + "value produced by a previous @jit compilation and cannot be " + "carried here: its backing IR lives in an already-finalized " + "compilation context. Pass it through the kernel's arguments " + "or recompute it in this compilation." + ) + + +def _pyir_promote_raw_scalar_carry(target_name: str, current_value: object) -> object: + """Serve a loop-region read of a raw scalar ``ir.Value`` write_arg + (``scf.while`` condition or loop-body entry) from its place cell, minting + the cell at the declared pre-region position on first sight (the raw-value + sibling of the wrapper mint in :func:`pyir_read`).""" + if pyir is None or _make_slot_key(target_name, None, None) is None: + return current_value + mv = _get_slot_mv(None, target_name) + if mv is None or not mv._is_ref_accessible(): + # A pre-region rebind may have D1-promoted this name already: adopt + # that row (body D1 stores keep writing it; a second cell would split + # body reads from body writes). + mv = _pyir_adopt_d1_carry_cell(target_name, current_value) + if mv is None or not mv._is_ref_accessible(): + _pre_ip = _pyir_promoted_region_pre_ip() + if _pre_ip is not None: + with _pre_ip: + mv = _create_ref(current_value) + else: + mv = _create_ref(current_value) + mv = _set_slot_mv(None, target_name, mv) + loaded = mv.load() + _attach_mutable_ref(loaded, mv, f"raw while-carry '{target_name}'") + return loaded + + +def _pyir_promoted_region_pre_ip() -> "ir.InsertionPoint | None": + """The declared pre-region position when a carry promotion fires with the + promoted loop op's own block as the ambient insertion point.""" + if ir is None: + return None + try: + owner = ir.InsertionPoint.current.block.owner + op = getattr(owner, "operation", owner) + if getattr(op, "name", None) in ("scf.while", "scf.for"): + return ir.InsertionPoint(op) + except Exception: + pass + return None + + +def _pyir_raw_dominates_region_exterior(raw: Any) -> bool: + """True when *raw* is defined directly in the function body region (its + definition's parent is the function op itself). Trace emission is + program-ordered, so a function-scope definition precedes -- and therefore + dominates -- every later trace point, including reads emitted after the + enclosing staged region closes. Conservative: False.""" + if ir is None: + return False + try: + owner = raw.owner + if isinstance(owner, ir.Block): + host = owner.owner # block argument: the op owning the block + else: + host = getattr(owner, "operation", owner).parent + host_op = getattr(host, "operation", host) + return _is_func_boundary_op(str(host_op.name)) + except Exception: + return False + + +def _promote_loop_carried_meta(target_name: str, current_value: object) -> object: + """Eagerly promote a read-before-write loop-carried meta primitive to a + function-entry ``pyir.ref`` and return a dominating load.""" + if pyir is None: + return current_value + d1_slot = _make_slot_key(target_name, None, None) + if d1_slot is None: + return current_value + initial_py = ( + current_value.python_value + if isinstance(current_value, _WatchedM) + else current_value + ) + if not isinstance(initial_py, (bool, int, float)): + return current_value + existing_ref = _slot_refs.get(d1_slot) + if existing_ref is None: + # Read-before-write => loop-carried: force the reset flag off. EXCEPTION: + # a strictly-shallower literal first-def is an outer-loop reset that must survive. + first_def_depth = _slot_first_def_depth.get(d1_slot) + outer_reset = ( + _slot_first_def_inside_cf.get(d1_slot, False) + and first_def_depth is not None + and first_def_depth < current_staged_cf_depth() + ) + if not outer_reset: + # Adopt the value-carried first-def facts (F-FIRSTDEF) only when + # the recorded reset position dominates this read; else plain carry. + adopted = False + fd_key = ( + current_value._slot_key + if isinstance(current_value, _WatchedM) + else None + ) + if fd_key is not None and fd_key != d1_slot: + fd_block = _slot_first_def_block.get(fd_key) + fd_depth = _slot_first_def_depth.get(fd_key) + try: + cur_block = ir.InsertionPoint.current.block + except Exception: + cur_block = None + if ( + _slot_first_def_inside_cf.get(fd_key, False) + and fd_block is not None + and fd_depth is not None + and _block_strictly_inside(cur_block, fd_block) + ): + _slot_first_def_inside_cf[d1_slot] = True + _slot_first_def_block[d1_slot] = fd_block + _slot_first_def_depth[d1_slot] = fd_depth + adopted = True + if not adopted: + _slot_first_def_inside_cf[d1_slot] = False + existing_ref = _meta_promote_slot(d1_slot, initial_py, target_name) + if existing_ref is None: + return current_value + return _load_as_dsl(existing_ref, place=d1_slot) + + +def _pyir_check_unpack_arity(seq: Any, expected: int, has_star: bool) -> None: + """Replicate CPython's unpack arity check for the decomposed unpack. + + The decomposition stores each element with ``tmp[i]``, which silently + ignores a length mismatch -- Python raises. Emitted once per sequence + target (each target of ``(p, q) = (c, r, s) = rhs`` checks its own arity + against the shared temp). + """ + got = len(seq) + if has_star: + if got < expected: + raise ValueError( + f"not enough values to unpack (expected at least {expected}, got {got})" + ) + return + if got < expected: + raise ValueError( + f"not enough values to unpack (expected {expected}, got {got})" + ) + if got > expected: + # CPython below 3.14 words this one without the ", got N" tail; the other + # two carried it already. + # Emit the 3.14 spelling everywhere because it is clear and the text + # does not depend on the interpreter the trace happens to run under. + raise ValueError(f"too many values to unpack (expected {expected}, got {got})") + + +def _pyir_pin_tuple_capture( + value: Any, _seen: "set[int] | None" = None, *, unpack_source: bool = False +) -> Any: + """Pin the watched-meta elements of a tuple-unpack RHS at the CAPTURE + position (called by the AST tuple-assign decomposition's temp statement). + + Python evaluates the whole RHS before any element store, so a swap must + consume each source's value AS OF the capture. + + A ``_WatchedM`` element materialised here leaves a position-recorded + constant; a later promotion rewrites it to a position-correct ``pyir.load``. + + Meta-ness is untouched; non-watched elements and non-container captures + pass through. + + ``unpack_source`` marks the pin that captures a whole unpack RHS (emitted by + ``_make_rhs_pin``); such a value is also MATERIALISED here when it is not + already an indexable sequence. The decomposition extracts elements by index + (``tmp[i]``) while Python's native unpack walks ``__iter__``, so a generator + / ``zip`` / set / dict source has to become a tuple exactly once (a + generator is single-use, and several chained targets share this one temp). + A non-iterable source raises ``TypeError`` from ``tuple()`` here, which is + what Python's own unpack does. + + The flag is OFF for the element pin the decomposition wraps around a + ``_pyir_nest_N`` extraction: that value is a single element, routinely a + scalar, and its own follow-up statement pins it as an unpack source later. + """ + if unpack_source and not isinstance(value, (tuple, list, str, bytes)): + value = tuple(value) + if pyir is None or not is_inside_staged_cf(): + return value + if isinstance(value, (tuple, list)): + # Recurse into NESTED tuple/list elements: their leaves defer to a + # follow-up statement of the unpack split, which runs AFTER the + # flat element stores -- an unpinned doubly-nested leaf would load + # the post-store value (swap pinning breaks one level down). The + # id set only breaks self-referential list cycles. + _seen = set() if _seen is None else _seen + if id(value) in _seen: + return value + _seen.add(id(value)) + pinned = [ + ( + _pyir_pin_tuple_capture(_elem, _seen) + if isinstance(_elem, (tuple, list)) + else _pyir_anchor_statement_read(_elem) + ) + for _elem in value + ] + if any(p is not e for p, e in zip(pinned, value)): + return ( + pinned + if isinstance(value, list) + else _rebuild_tuple_like(value, pinned) + ) + return value + return _pyir_anchor_statement_read(value) + + +def _pyir_anchor_statement_read(value: Any) -> Any: + """Anchor one managed-place read at ITS OWN program point (statement-top + read-anchor rule, the scalar sibling of :func:`_pyir_pin_tuple_capture`). + + Emitted by the preprocessor for a read Python evaluates BEFORE a + same-statement store of the same place (a pre-walrus read). + + * ``_WatchedM``: materialise the constant HERE so a later promotion + rewrites it to a ``pyir.load`` preceding the store; meta-ness untouched. + + * cell-paired staged wrapper: return a fresh SNAPSHOT ``mv.load()`` (no + ``_mutable_ref``), so later reads never follow the cell past the store. + + Everything else passes through untouched. + """ + if pyir is None or not is_inside_staged_cf(): + return value + if isinstance(value, _WatchedM): + # PROMOTED place: the read must be a load AT THIS POSITION (a deferred + # materialisation could land after the same-statement store). + slot = getattr(value, "_slot_key", None) + if slot is not None: + _ref = _slot_refs.get(slot) + if isinstance(_ref, ir.Value) and _value_dominates_current_ip(_ref): + try: + return _load_as_dsl(_ref, attach=False, place=slot) + except Exception: + pass + try: + value.ir_value() + except Exception: + pass # no context / unsupported width -- consumer decides + return value + mv = getattr(value, "_mutable_ref", None) + if mv is not None: + try: + if mv.ref is not None and mv._is_ref_accessible(): + return mv.load() + except Exception: + pass + return value + + +def _pyir_watched_dict_read(container: "_WatchedDict", key: Any, val: Any) -> Any: + """Object-side READ choke for one adopted dict entry: the owner-keyed + ``pyir_read`` lifecycle from ANY Python code; nested plain dicts adopt. + An instance-``__dict__`` mapping routes str-keyed items to the INSTANCE's + attr row -- one place row per storage, whatever the spelling.""" + label = _pyir_watched_dict_label(container, key) + if type(val) is dict or isinstance(val, _WatchedDict): + return _pyir_adopt_dict_value(container, key, val, label=label) + if type(val) is list or isinstance(val, _WatchedList): + return _pyir_adopt_list_value(container, key, val, label=label) + # A plain OBJECT surfacing through a container leg: its attr legs are + # places one composition deeper (root -> dict-leg -> attr-leg). + _pyir_register_container_held_object(val) + place_owner = _pyir_instance_dict_owner(container, key) + return pyir_read( + label, + val, + attach_ref=False, + owner=place_owner if place_owner is not None else container, + slot_name=key, + ) + + +def _pyir_watched_dict_write( + container: "_WatchedDict", + key: Any, + value: Any, + filename: str, + lineno: int, +) -> Any: + """Object-side WRITE choke for one adopted dict entry: the AST write choke's + read+assign lifecycle; unpromoted meta changes in dynamic staged CF are refused. + An instance-``__dict__`` mapping routes str-keyed items to the INSTANCE's + attr row -- one place row per storage, whatever the spelling.""" + label = _pyir_watched_dict_label(container, key) + # A key must NAME a compile-time place: staged / IR-backed keys are refused; + # any hashable trace-time value is a valid place name. + if _is_staged_value(key) or isinstance(key, ir.Value): + raise DSLUserCodeError( + DiagId.CONTAINER_DICT_KEY_STAGED, + filename=filename, + lineno=lineno, + var=getattr(container, "_pyir_label", None) or "dict", + ) + # The same wall one hop deeper: a plain-object key whose own ``__hash__`` + # consumes a staged value launders runtime identity into the key. Probe + # the key's hash under the trace-scoped staged-hash witness; only this + # write-site correlation refuses -- hashing outside a keyed write never + # does. An unhashable key raises the plain TypeError, as any dict write. + _hash_mark = _pyir_staged_hash_witness_count() + hash(key) + if _pyir_staged_hash_witness_count() != _hash_mark: + raise DSLUserCodeError( + DiagId.CONTAINER_DICT_KEY_STAGED_HASH, + filename=filename, + lineno=lineno, + var=getattr(container, "_pyir_label", None) or "dict", + ) + _inst = _pyir_instance_dict_owner(container, key) + place_owner = _inst if _inst is not None else container + missing = not dict.__contains__(container, key) + if missing: + # New-key insertion through SUBSCRIPT syntax stays supported (first-def + # mints the cell; stale reads refuse at the READ). Mutator methods stay loud. + old: Any = None + # Inside dynamic staged CF the key SET changed on this one traced pass + # only: record it so a later whole-key-set consumption (membership, + # lookup miss) refuses instead of baking this pass's truth. + if is_inside_staged_cf() and not is_inside_constexpr_loop(): + _PYIR_DICT_CF_CREATED_KEYS.setdefault(id(container), (container, set()))[ + 1 + ].add(key) + # Gate-once evidence: a first-def traced under a folded constant ``if`` arm + # lifts keep-meta when the arm provably latched the gate (see pyir_state). + if _PYIR_FOLD_FIRSTDEF_STACK: + try: + _fd_slot = _make_slot_key(None, place_owner, key) + if _fd_slot is not None: + _PYIR_FOLD_FIRSTDEF_STACK[-1].append(_fd_slot) + except Exception: + pass + else: + old = pyir_read( + label, + dict.__getitem__(container, key), + owner=place_owner, + slot_name=key, + ) + if type(value) is dict: + value = _pyir_adopt_dict_value(container, key, value, label=label) + elif type(value) is list: + value = _pyir_adopt_list_value(container, key, value, label=label) + result = pyir_assign( + label, old, value, filename, lineno, owner=place_owner, slot_name=key + ) + # A first-def inside a PLAIN callee within a staged loop has unprovable + # multiplicity: lift keep-meta so later meta changes fail-close loudly. + if ( + missing + and _PYIR_BOUNDARY_CALLEE_DEPTH[0] > 0 + and not _PYIR_FOLD_FIRSTDEF_STACK + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + and _innermost_enclosing_loop_op_at_ip() is not None + ): + try: + _fd_slot = _make_slot_key(None, place_owner, key) + if _fd_slot is not None: + _slot_first_def_inside_cf[_fd_slot] = False + except Exception: + pass + # A NEW-key insert inside a constexpr-unrolled loop pins foreign-bound staged + # values at insertion-time SSA (the staged-container-insert freeze). + if missing and is_inside_constexpr_loop() and _carries_foreign_slot_binding(result): + result = _freeze_foreign_slot_binding( + result, "constexpr-loop watched-dict insert" + ) + # Mark the dict as carrying tracked writes: drives get()-miss loudness (a + # miss default would bake fixed while sibling entries update). + if _is_staged_value(result) or is_inside_staged_cf(): + try: + container._pyir_staged_writes = True + except Exception: + pass + # Loud residue: a meta->meta change on an un-promoted leg in dynamic staged CF + # refuses; declared keep-meta forms (constexpr, reset, unchanged, first def) pass. + if ( + old is not None + and isinstance(old, (bool, int, float)) + and isinstance(result, (bool, int, float)) + and not _is_staged_value(result) + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + ): + _d1 = _make_slot_key(None, place_owner, key) + _old_p = old.python_value if isinstance(old, _WatchedM) else old + _new_p = result.python_value if isinstance(result, _WatchedM) else result + if ( + _d1 is not None + and _d1 not in _slot_refs + and not _slot_first_def_inside_cf.get(_d1, False) + and _old_p != _new_p + ): raise DSLUserCodeError( - DiagId.SCOPE_CTXMGR_TRACE_ONLY, - ctx_type=tp.__name__, + DiagId.CONTAINER_DICT_META_WRITE_UNPROMOTED, + filename=filename, + lineno=lineno, + var=label, + old_value=repr(_old_p), + new_value=repr(_new_p), ) - return cm - for dunder in ("__enter__", "__exit__"): - fn = getattr(tp, dunder, None) - if fn is None: - continue - code = getattr(fn, "__code__", None) - dunder_is_internal = is_dsl_internal_code( - getattr(code, "co_filename", None), getattr(fn, "__module__", "") or "" + return result + + +def _pyir_watched_dict_mutate( + container: "_WatchedDict", op_name: str, detail: str +) -> None: + """Object-side STRUCTURAL-mutation choke (key-set changes): loud inside dynamic + staged CF only; constexpr scopes realize the mutation and route the insert freeze.""" + _pyir_spec_structural_amend(container) + if op_name in ("setdefault", "update"): + _pyir_freeze_staged_container_inserts(container, "update") + if is_inside_staged_cf() and not is_inside_constexpr_loop(): + raise DSLUserCodeError( + DiagId.CONTAINER_DICT_KEY_SET_MUTATED, + var=getattr(container, "_pyir_label", None) or "dict", + detail=detail, ) - if _is_dsl_traced_callable(fn) or dunder_is_internal: - continue + + +def _pyir_watched_dict_get_miss(container: "_WatchedDict", key: Any) -> None: + """Object-side ``get()``-MISS choke: a miss default on a dict with tracked + writes inside dynamic staged CF would bake a leaked-arm constant -- loud.""" + if ( + getattr(container, "_pyir_staged_writes", False) + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + ): + raise DSLUserCodeError( + DiagId.CONTAINER_DICT_GET_MISS_IN_STAGED_CF, + var=getattr(container, "_pyir_label", None) or "dict", + detail=repr(key), + ) + + +def _pyir_watched_dict_iterate(container: Any, method: str) -> Any: + """ITERATION choke for tracked dicts inside dynamic staged CF. The key + set is trace-time STABLE (structural mutation refuses at the mutator + choke), so enumerating it is exact -- but the raw stored VALUES are a + trace-time snapshot that bypasses the per-entry read choke, so + ``.values()``/``.items()`` serve every entry through the same tracked + read path ``d[k]`` uses (the carried slot read); a write driven by the + pair then stores through the tracked path naturally. Key walks (bare + ``iter()``/``.keys()``) raw-delegate: per-slot reads re-enter the item + choke themselves. The one UNSTABLE case refuses: a key CREATED inside + dynamic staged CF (subscript first-def) exists on the one traced pass + only, so the enumerated key set cannot follow the runtime path. + Returns the served sequence, or None to raw-delegate (constexpr scopes + realize the walk at trace time).""" + if not is_inside_staged_cf() or is_inside_constexpr_loop(): + return None + entry = _PYIR_DICT_CF_CREATED_KEYS.get(id(container)) + if entry is not None: raise DSLUserCodeError( - DiagId.SCOPE_CTXMGR_TRACE_ONLY, - ctx_type=tp.__name__, + DiagId.CONTAINER_DICT_ITERATED_IN_STAGED_CF, + var=getattr(container, "_pyir_label", None) or "dict", + method=method, + detail=", ".join(sorted(repr(k) for k in entry[1])), ) - return cm + name = method.strip(".()") + if name not in ("values", "items"): + return None + served = [] + for key in list(dict.keys(container)): + val = container[key] + served.append(val if name == "values" else (key, val)) + return served -def _pyir_while_cond(cond: object) -> object: - """Lift a ``_WatchedM`` while condition to its staged DSL Numeric. +def _pyir_watched_list_index( + container: "_WatchedList", key: Any, label: str +) -> "int | None": + """Validate and normalise a list-leg index: staged indices refuse, negative + trace-time indices normalise, anything else returns ``None`` (raw delegation).""" + if _is_staged_value(key) or isinstance(key, ir.Value): + raise DSLUserCodeError( + DiagId.CONTAINER_LIST_INDEX_STAGED, + var=getattr(container, "_pyir_label", None) or "list", + ) + if not isinstance(key, int): + return None + idx = int(key) + if idx < 0: + idx += list.__len__(container) + if idx < 0 or idx >= list.__len__(container): + return None + return idx - The while executor truth-tests the before-block's condition; a watched - fold bakes ``scf.condition(true)`` -- an unkillable runtime hang once - the body mutates the slot. Lift when a connected slot is genuinely - promoted (in ``_slot_refs``) or in the region's syntactic write-set - (tagged up front by ``pyir_tag_pending_writes`` in the before-block - prologue); otherwise keep the fold with a forced witness so a later - store still raises the two-location stale diagnostic instead of hanging. - NOTE (master bridge): the donor keyed "already promoted" on - ``_slot_stored`` (D1v2, absent here); master tracks promotion as - membership in ``_slot_refs`` (set in ``_meta_promote_slot``), so the - gate uses ``_slot_refs.keys()`` in its place. - """ - from .pyir_core import ( - _WatchedM, - _record_fold_witness, - _watched_to_dsl, - _wrapper_slots, +def _pyir_watched_list_read(container: "_WatchedList", key: Any) -> Any: + """Object-side READ choke for one adopted list element (owner-keyed snapshot + read; nested plain containers adopt; slices never reach here).""" + label = _pyir_watched_list_label(container, key) + idx = _pyir_watched_list_index(container, key, label) + if idx is None: + return list.__getitem__(container, key) + val = list.__getitem__(container, idx) + if type(val) is dict or isinstance(val, _WatchedDict): + return _pyir_adopt_dict_value(container, idx, val, label=label) + if type(val) is list or isinstance(val, _WatchedList): + return _pyir_adopt_list_value(container, idx, val, label=label) + # A plain OBJECT surfacing through a container leg: its attr legs are + # places one composition deeper (root -> list-leg -> attr-leg). + _pyir_register_container_held_object(val) + return pyir_read( + label, + val, + attach_ref=False, + owner=container, + slot_name=idx, + ) + + +def _pyir_watched_list_write( + container: "_WatchedList", + key: Any, + value: Any, + filename: str, + lineno: int, +) -> Any: + """Object-side WRITE choke for one adopted list element: the dict write choke's + lifecycle on integer legs (no first-def arm; unpromoted meta changes refuse).""" + label = _pyir_watched_list_label(container, key) + idx = _pyir_watched_list_index(container, key, label) + if idx is None: + # Wrong-type / out-of-range index: the caller's raw store raises the + # plain-Python error. + return value + old = pyir_read( + label, + list.__getitem__(container, idx), + owner=container, + slot_name=idx, ) - from .pyir_state import ( - _slot_pending_store, - _slot_refs, + if type(value) is dict: + value = _pyir_adopt_dict_value(container, idx, value, label=label) + elif type(value) is list: + value = _pyir_adopt_list_value(container, idx, value, label=label) + result = pyir_assign( + label, old, value, filename, lineno, owner=container, slot_name=idx ) + if ( + old is not None + and isinstance(old, (bool, int, float)) + and isinstance(result, (bool, int, float)) + and not _is_staged_value(result) + and is_inside_staged_cf() + and not is_inside_constexpr_loop() + ): + _d1 = _make_slot_key(None, container, idx) + _old_p = old.python_value if isinstance(old, _WatchedM) else old + _new_p = result.python_value if isinstance(result, _WatchedM) else result + if ( + _d1 is not None + and _d1 not in _slot_refs + and not _slot_first_def_inside_cf.get(_d1, False) + and _old_p != _new_p + ): + raise DSLUserCodeError( + DiagId.CONTAINER_LIST_META_WRITE_UNPROMOTED, + filename=filename, + lineno=lineno, + var=label, + old_value=repr(_old_p), + new_value=repr(_new_p), + ) + return result - if not is_inside_staged_cf() or not isinstance(cond, _WatchedM): - return cond - if not (_wrapper_slots(cond) & (_slot_refs.keys() | _slot_pending_store)): - _record_fold_witness( - cond, - None, - cond.python_value, - "a Python `if`/`while`/`not` test", - force_consumer=True, + +def _pyir_watched_list_mutate( + container: "_WatchedList", op_name: str, detail: str +) -> None: + """Object-side STRUCTURAL-mutation choke (length/order changes): loud inside + dynamic staged CF; constexpr scopes realize the mutation and route the freeze.""" + _pyir_spec_structural_amend(container) + if op_name in _PYIR_LIST_INSERT_MUTATORS: + _pyir_freeze_staged_container_inserts(container, op_name) + if is_inside_staged_cf() and not is_inside_constexpr_loop(): + raise DSLUserCodeError( + DiagId.CONTAINER_LIST_SHAPE_MUTATED, + var=getattr(container, "_pyir_label", None) or "list", + detail=detail, ) - return cond - try: - return _watched_to_dsl(cond) - except Exception: # noqa: BLE001 — keep the fold on lift failure - return cond -__all__ = [name for name in list(globals()) if not name.startswith("__")] +# Install the watched-container choke hooks (the classes live in the core +# layer, which cannot import this one; see _WATCHED_*_HOOK in pyir_state). +_WATCHED_DICT_READ_HOOK[0] = _pyir_watched_dict_read +_WATCHED_DICT_WRITE_HOOK[0] = _pyir_watched_dict_write +_WATCHED_DICT_MUTATOR_HOOK[0] = _pyir_watched_dict_mutate +_WATCHED_DICT_GET_MISS_HOOK[0] = _pyir_watched_dict_get_miss +_WATCHED_DICT_ITER_HOOK[0] = _pyir_watched_dict_iterate +_WATCHED_LIST_READ_HOOK[0] = _pyir_watched_list_read +_WATCHED_LIST_WRITE_HOOK[0] = _pyir_watched_list_write +_WATCHED_LIST_MUTATOR_HOOK[0] = _pyir_watched_list_mutate + + +# Static export surface (regenerate with scripts/gen_pyir_all.py): the +# module's OWN names -- its namespace minus what lower chain modules +# already export and minus the generator's INTERNAL_ONLY registry. +__all__ = [ + "_decompose_tuple", + "_unalias_tuple_leaves", + "_pyir_check_no_complex_m2m_call", + "_pyir_freeze_staged_container_inserts", + "pyir_promote_loop_body_arg", + "_pyir_record_fresh_entry_birth", + "pyir_assign", + "PYIR_REGION_ATTR_WRITES_ATTR", + "PYIR_REGION_METHOD_CALLS_ATTR", + "PYIR_REGION_FREE_CALLS_ATTR", + "pyir_tag_region_attr_writes", + "pyir_note_attr_first_def", + "pyir_generation_probe", + "_pyir_judge_fabricated_attr_read", + "_PYIR_FGET_DEREF_OPS", + "pyir_read", + "pyir_bind_param", + "_pyir_obs_read", + "_pyir_obs_global_read", + "_pyir_post_subscript_read", + "_pyir_pre_subscript_assign", + "_attr_owner_plain_storage_setattr", + "_pyir_pre_attr_assign", + "_pyir_post_attr_assign", + "pyir_promote_while_carried_arg", + "_pyir_check_unpack_arity", + "_pyir_pin_tuple_capture", + "_pyir_anchor_statement_read", +] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_loop_carry.py b/python/CuTeDSL/cutlass/base_dsl/pyir_loop_carry.py new file mode 100644 index 0000000000..eef2980f1b --- /dev/null +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_loop_carry.py @@ -0,0 +1,1135 @@ +# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: LicenseRef-NvidiaProprietary +# +# Use of this software is governed by the terms and conditions of the +# NVIDIA End User License Agreement (EULA), available at: +# https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/license.html +# +# Any use, reproduction, disclosure, or distribution of this software +# and related documentation outside the scope permitted by the EULA +# is strictly prohibited. + + +"""PyIR runtime -- loop-carry layer; see facade for the public surface. + +Loop carries exist because a value written in a loop body must feed the next +iteration before the region closes. Staged writes inside scf.if regions need +no carry here: they store into the variable's function-entry pyir.ref slot and +the mem2reg lowering in convert-pyir-to-scf materializes them as scf.if +results; only the opaque-leaf companion _pyir_carry_if_region_mutated_leaves +lives in this module. +""" + +from .pyir_state import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_core import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_corewalk import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) + + +def _pyir_push_loop_body_scope(body_block: "ir.Block") -> None: + """Declare *body_block* as the innermost open staged-loop body (F-BIRTHPOS); + a body-region entry is a region-epoch boundary (F-GEN).""" + _pyir_bump_region_epoch() + _pyir_region_entry_push() + _pyir_open_loop_body_blocks.append(body_block) + + +def _pyir_pop_loop_body_scope() -> None: + """Close the innermost open loop-body declaration (a region-epoch boundary).""" + _pyir_bump_region_epoch() + _pyir_region_entry_pop() + if _pyir_open_loop_body_blocks: + _pyir_open_loop_body_blocks.pop() + + +def _meta_promote_slot( + slot_key: Any, + initial_py_value: Any, + target_name: str | None = None, + filename: str | None = None, + lineno: int | None = None, + promoted_value: Any = None, + display_name: str | None = None, +) -> "ir.Value | None": + """Promote a slot to D1 tracking (AUTO_M2S only; refuses when the flag + is off -- ``PHASE_AUTO_PROMOTE_DISABLED``). + + Creates a ``pyir.ref`` at the enclosing function's entry block + initialized with *initial_py_value*, then walks every previously + baked ``arith.constant`` recorded under *slot_key* and replaces its + uses with a freshly-emitted ``pyir.load %ref``. The replaced + constants become dead and DCE cleans them up; downstream arith ops + pick up the load via SSA edges automatically -- no per-derivation + rewriting needed. + + Warns which variable was promoted (value + file:line of the mutation), + with the same wording as the Mp→S warning in ``pyir_read``. + + Returns the new ``pyir.ref`` SSA value, or ``None`` if there is no + enclosing function (e.g. tracing happens outside an MLIR context). + """ + if pyir is None: + return None + entry_block = _get_function_entry_block() + if entry_block is None: + return None + + # Implicit Meta-to-Staged promotion is an AUTO_M2S feature; with the flag + # off the mode's contract is "never silently wrong", so every promotion + # request refuses loudly instead of rewriting the slot. + if not is_auto_m2s_enabled(): + _promoted_cls = _declared_m2s_promotion_class(initial_py_value) + raise DSLUserCodeError( + DiagId.PHASE_AUTO_PROMOTE_DISABLED, + filename=filename, + lineno=lineno, + var=target_name or display_name or str(slot_key), + value=repr(initial_py_value), + type=_promoted_cls.__name__ if _promoted_cls is not None else "Int32", + ) + + # This place gated a plain callee's literal store that boundary commit turned + # into a per-iteration constant; staging it breaks that proof -- refuse loudly. + _guard_read = _PYIR_BOUNDARY_FLIP_GUARD_READS.get(slot_key) + if _guard_read is not None: + _gr_name, _gr_flipped, _gr_file, _gr_line = _guard_read + raise DSLUserCodeError( + DiagId.BOUNDARY_FLIP_GUARD_STAGED, + filename=filename, + lineno=lineno, + name=_gr_name, + flipped=_gr_flipped, + def_file=_gr_file, + def_line=_gr_line, + ) + + # This place was read through a closure cell by a plain callee inside staged + # CF; no rewrite can retarget that baked read -- refuse loudly. + _cell_read = _pyir_boundary_cell_read_for_local(slot_key) + if _cell_read is not None: + _cr_name, _cr_file, _cr_line = _cell_read + raise DSLUserCodeError( + DiagId.BOUNDARY_CLOSURE_READ_THEN_PROMOTED, + filename=filename, + lineno=lineno, + name=_cr_name, + def_file=_cr_file, + def_line=_cr_line, + ) + + # A place already consumed as trace-time structure has no retargetable SSA + # constant, so promotion refuses -- unless the promotion seed AND the + # triggering write (when there is one) provably re-establish the consumed + # value, in which case the baked structure stays valid on every iteration. + # A write that stays in the region birthing the place re-runs that binding + # every iteration, so the bake stays valid and promotion is safe even when + # the triggering write is staged. Only a place born above the loop carries. + _sc = _PYIR_STRUCTURAL_META_CONSUMPTIONS.get(slot_key) + if _sc is not None and _pyir_structural_bake_is_reseeded(slot_key): + _sc = None + if _sc is not None and ( + _pyir_structural_value_conflicts(_sc, initial_py_value) + or ( + promoted_value is not None + and _pyir_structural_value_conflicts(_sc, promoted_value) + ) + ): + _sc_value, _sc_file, _sc_line = _sc + raise DSLUserCodeError( + DiagId.PHASE_STRUCTURAL_CONSTANT_MUTATED, + filename=filename, + lineno=lineno, + var=target_name or str(slot_key), + value=_sc_value, + read_file=_sc_file, + read_line=_sc_line, + ) + + # F-SPEC: the promotion bakes the meta payload as the cell's seed -- a + # specialization fact of this trace when the place is root-pathable. + _pyir_spec_record_read(slot_key, initial_py_value) + + # A promotion triggered by a staged write must type the ref pointee (and every + # reset constant) by the staged value's type, else a later store fails verification. + _staged_target_type = None + if promoted_value is not None: + from .multi_stage_manager import _is_staged_value + + if _is_staged_value(promoted_value): + _staged_target_type = type(promoted_value) + + def _emit_init() -> "ir.Value": + if _staged_target_type is not None: + try: + return _staged_target_type(initial_py_value).ir_value() + except (TypeError, ValueError, AttributeError): + pass + return _emit_constant_at_current_ip(initial_py_value) + + # F-BIRTHPOS: a region-born place seeds itself with a real store at its + # recorded birth block (per-iteration in a loop body, per-branch in an if + # arm) and takes an undefined-placeholder entry init, so a path where the + # binding never ran is caught by the used-placeholder scan. + _birth_block = ( + _slot_first_def_block.get(slot_key) + if _slot_first_def_inside_cf.get(slot_key, False) + else None + ) + if _birth_block is not None: + with ir.InsertionPoint.at_block_begin(_birth_block): + seed_ir = _emit_init() + with ir.InsertionPoint.at_block_begin(entry_block): + init_ir = _make_raw_placeholder_init(seed_ir.type) + if init_ir is None: + init_ir = _emit_init() + ref = pyir.ref(init_ir) + with ir.InsertionPoint.after(_get_defining_operation(seed_ir)): + _pyir_emit_store(seed_ir, ref) + else: + with ir.InsertionPoint.at_block_begin(entry_block): + initial_ir = _emit_init() + ref = pyir.ref(initial_ir) + _slot_refs[slot_key] = ref + # F-TYPEID: the promotion declares the row's wrapper class; reads + # reconstruct through it (a later staged store advances it). + _pyir_record_slot_template( + slot_key, + _pyir_declared_promotion_template(initial_py_value, _staged_target_type), + ) + + # User-visible warning: a Python value just became staged-tracked; show target, value, + # file:line so users can hoist the init out of CF. + if target_name is not None: + promoted_cls = _declared_m2s_promotion_class(initial_py_value) + promoted_type = ( + promoted_cls.__name__ + if promoted_cls is not None + else type(initial_py_value).__name__ + ) + # Stale-local risk: the slot is now a ``pyir.ref`` but the caller's Python binding still + # holds the constexpr, so reads OUTSIDE this region fold the stale value. + from .diagnostics import WarnId, report_warning + + report_warning( + WarnId.PHASE_AUTO_PROMOTED_TO_STAGED, + filename=filename, + lineno=lineno, + stacklevel=4, + var=target_name, + value=repr(initial_py_value), + type=promoted_type, + ) + + # In-region first-def with no recorded birth block: seed before the earliest + # baked use, which the follow-on rewrite turns into a dominated ``pyir.load``. + if _slot_first_def_inside_cf.get(slot_key, False) and _birth_block is None: + try: + uses = _meta_uses.get(slot_key, []) + if uses and not isinstance(uses[0].owner, ir.Block): + owner_op = uses[0].owner + defining_op = ( + owner_op + if isinstance(owner_op, ir.Operation) + else getattr(owner_op, "operation", owner_op) + ) + with ir.InsertionPoint(defining_op): + reset_ir = _emit_init() + _pyir_emit_store(reset_ir, ref) + except Exception as exc: + log().info( + "[_meta_promote_slot] %s: reset-store insert failed: %s", + slot_key, + exc, + ) + + # Re-materialise idempotent writes the promotion gates skipped: each skip + # left an anchor constant at its write site, and Python DID execute the + # assignment there -- once the slot is a cell, every such site must reset + # the cell (e.g. a body-top ``x = 128`` re-arms 128 each iteration even + # after a later conditional write stored a different value). Idempotence + # guarantees every anchor's payload equals ``initial_py_value``, but the + # anchor's own IR type may not match the ref pointee when the promotion + # was triggered by a staged write (``_staged_target_type``), so the anchor + # serves as the insertion position only and the stored value is re-emitted + # through ``_emit_init`` at the ref's type; the anchor itself goes dead. + for _anchor in _meta_idempotent_write_anchors.pop(slot_key, []): + with ir.InsertionPoint.after(_get_defining_operation(_anchor)): + _pyir_emit_store(_emit_init(), ref) + + # Rewrite baked uses to ``pyir.load %ref``, but ONLY those whose baked value + # matches the ref init: an earlier-value use keeps its literal (else corrupted). + _rewrite_recorded_meta_uses(slot_key, ref, initial_py_value) + + return ref + + +def _rewrite_recorded_meta_uses( + slot_key: Any, ref: "ir.Value", initial_py_value: Any +) -> None: + """Replace each baked constant under *slot_key* whose value matches + *initial_py_value* with a position-correct ``pyir.load`` of *ref*.""" + for const_val in _meta_uses.pop(slot_key, []): + try: + owner = const_val.owner + if isinstance(owner, ir.Block): + continue # block argument has no "before" insertion point + baked_value = _const_value_of(const_val) + if baked_value is not _NO_CONST_VALUE and not _const_values_equal( + baked_value, initial_py_value + ): + log().info( + "[_rewrite_recorded_meta_uses] %s: keep literal %r (init=%r) -- " + "earlier-value baked use, not carried through ref", + slot_key, + baked_value, + initial_py_value, + ) + continue + defining_op = ( + owner + if isinstance(owner, ir.Operation) + else getattr(owner, "operation", owner) + ) + with ir.InsertionPoint(defining_op): + loaded = pyir.load(ref) + if not _replace_value_uses(const_val, loaded): + # A missing RAUW binding method is an environment defect that + # silently degrades promotion (stale-literal reads); warn, not info. + log().warning( + "[_rewrite_recorded_meta_uses] %s: replace_all_uses_with " + "missing on ir.Value", + slot_key, + ) + else: + # Record the rewrite so a cached materialisation of this constant + # hands new consumers the load instead of the dead constant. + try: + _META_CONST_REPLACEMENTS[const_val] = loaded + except TypeError: + pass + except Exception as exc: + log().info( + "[_rewrite_recorded_meta_uses] %s: replace failed: %s", slot_key, exc + ) + + +def _stage_compound_value(value: Any, seen: "set[int]") -> "tuple[Any, bool] | None": + """Recursive worker for :func:`_stage_meta_compound_leaves`. + + Returns ``(staged_value, changed)``, or ``None`` when some leaf + cannot be staged. Compounds are rebuilt, never mutated in place: + a pre-loop alias must keep the untouched original, as it would in + plain Python.""" + from .multi_stage_manager import _is_staged_value + from .pyir_call_boundary import _pyir_boundary_module_is_user + + if value is None or type(value) in (str, bytes): + return (value, False) + py = _pyir_unwrap_meta_primitive(value) + if py is not None and type(py) in (bool, int, float): + promoted = _auto_promote_primitive(py) + return None if promoted is None else (promoted, True) + if _is_staged_value(value): + return (value, False) # already staged: the carry engine handles it + if _implements_dynamic_expression(value): + # A value-protocol object carries through its own extract/reconstruct + # pair: leave it alone, as the all-meta gate does. + return (value, False) + if id(value) in seen: + return None # cycle: no sound rebuild order + if type(value) in (tuple, list, _WatchedList): + if not value: + return (value, False) # empty: fixed structure, keep as-is + seen.add(id(value)) + elems: "list[Any]" = [] + changed = False + for elem in list.__iter__(value) if isinstance(value, list) else value: + sub = _stage_compound_value(elem, seen) + if sub is None: + return None + elems.append(sub[0]) + changed = changed or sub[1] + if not changed: + return (value, False) + if isinstance(value, list): + return (elems, True) # staged replica is a plain list + return (_rebuild_tuple_like(value, elems), True) + if type(value) in (dict, _WatchedDict): + if not value: + return (value, False) + seen.add(id(value)) + entries: "dict[Any, Any]" = {} + changed = False + for key, elem in dict.items(value): + if _is_staged_value(key): + return None # keys are compile-time structure + sub = _stage_compound_value(elem, seen) + if sub is None: + return None + entries[key] = sub[0] + changed = changed or sub[1] + if not changed: + return (value, False) + return (entries, True) + attrs = getattr(value, "__dict__", None) + if ( + isinstance(attrs, dict) + and attrs + and _pyir_boundary_module_is_user(type(value).__module__) + ): + seen.add(id(value)) + staged_attrs: "dict[str, Any]" = {} + changed = False + for name, elem in attrs.items(): + sub = _stage_compound_value(elem, seen) + if sub is None: + return None + staged_attrs[name] = sub[0] + changed = changed or sub[1] + if not changed: + return (value, False) + try: + clone = object.__new__(type(value)) + except Exception: + return None # __new__ rejected: cannot clone this class + for name, elem in staged_attrs.items(): + _pyir_setattr_raw(clone, name, elem) + return (clone, True) + return None + + +def _stage_meta_compound_leaves(value: Any) -> "Any | None": + """Rebuild a meta compound (tuple, list, dict, or user-class object) + with every bool/int/float leaf staged, so the per-leaf carry engine + can thread the loop body's rebuilds. Returns ``None`` when a leaf + cannot be staged or nothing needed staging. Tuple subclasses + (namedtuple) are excluded: element reads spelled ``t.field`` bypass + the element cells and would silently keep stale values.""" + result = _stage_compound_value(value, set()) + if result is None or not result[1]: + return None + return result[0] + + +def _pyir_snapshot_region_arg(arg: Any) -> "list[tuple[Any, str, ir.Value]] | None": + """Snapshot *arg*'s in-place ``ir.Value`` leaf holders before a region body.""" + holders = _pyir_walk_ir_value_holders(arg) + if holders: + return holders + # A TOP-LEVEL tuple/list of value-trees yields nothing from the value-tree + # walk; snapshot the CONCATENATED element leaves with the tuple as holder. + if _pyir_is_carryable_tuple(arg): + leaves = _pyir_extract_any_leaf_values(arg) + if leaves: + return [(arg, _PYIR_TUPLE_LEAF_ATTR, lv) for lv in leaves] + return _self_leaf_snapshot(arg) + + +# Region-kind wording for the shared carry core: (type-mismatch phrase, +# fail-loud wall label, carried-as phrase). +_CARRY_PHRASES = { + "loop": ("re-derived per iteration", "loop leaf carry", "iter_arg"), + "if": ("re-derived per branch", "scf.if leaf carry", "scf.if result"), +} + + +def _pyir_carry_region_mutated_leaves( + arg: Any, + snapshot: "list[tuple[Any, str, ir.Value]] | None", + region_op: "ir.Operation", + region_block: "ir.Block", + context: str, + *, + kind: str, + opaque_only: bool, + skip_inner_loop_carried: bool, + except_users_outside_region: bool, + self_leaf_arg_index_only: bool, + ref_cache: "dict[int, ir.Value] | None" = None, + exclude_aliased_leaves: "list[ir.Value] | None" = None, + arg_index: int = -1, + arg_name: "str | None" = None, +) -> "list[tuple[Any, str, ir.Value, Any, int, int]]": + """Shared core of the loop/if region-carry twins: pair each pre-region + snapshot leaf with its current value, and carry every genuinely trapped + update through a ``pyir.ref`` (region-entry load, region-scoped RAUW, + region-exit store) so the C++ pass lifts it to an iter_arg / region + result. The wrappers fix the policy flags; every flag encodes a + load-bearing semantic difference between the two regions -- see each + wrapper's docstring.""" + phrase_mismatch, phrase_wall, phrase_carried = _CARRY_PHRASES[kind] + records: "list[tuple[Any, Any, ir.Value, Any, int, int]]" = [] + if pyir is None or snapshot is None: + return records + # Pair each snapshot leaf with the region's current value (in-place: ``rebind_holder`` set; whole-object + # rebind: ``reconstruct_obj`` + ``leaf_index`` let the caller reconstruct post-region from leaf loads). + for ( + rebind_holder, + attr_name, + snap_value, + cur_value, + reconstruct_obj, + leaf_index, + ) in _pyir_resolve_loop_leaf_updates(arg, snapshot): + # Trigger only on a genuine escaped update: snapshot defined OUTSIDE the region (dominates -> ref + # anchor), new value trapped INSIDE -- it must carry forward here, not be reverted at outer scope. + if not isinstance(cur_value, ir.Value) or cur_value is snap_value: + continue + if not isinstance(snap_value, ir.Value): + continue + if opaque_only and not _is_opaque_leaf_type(snap_value): + # A scalar leaf is already carried by the M2S function-entry ref; + # only an opaque leaf needs region-level carry. Avoid double-carry. + continue + if skip_inner_loop_carried and _pyir_value_is_inner_loop_carried( + cur_value, region_op + ): + # ``scf.while``: this leaf is the post-close load of a ref a NESTED loop already carried; that + # loop owns the carry. Avoid double-carry. + continue + # ALIASED-leaf exclusion (standalone ``scf.if``): skip a leaf whose pre-``if`` SSA is also held by a + # captured free variable -- yielding it as an ``scf.if`` result would point aliases at a trapped load. + if exclude_aliased_leaves is not None and any( + _same_ir_value(snap_value, _al) for _al in exclude_aliased_leaves + ): + continue + if _ir_value_defined_inside_op(snap_value, region_op): + # Snapshot defined inside the region cannot anchor a pre-region ref (defensive: the snapshot + # is taken at the outer scope). + continue + if not _ir_value_defined_inside_op(cur_value, region_op): + # The mutation is a legitimate forward update defined outside the region -- nothing trapped. + continue + if cur_value.type != snap_value.type: + # An scf iter_arg / region result is type-invariant by construction, so a leaf + # rebound to a DIFFERENTLY-TYPED value cannot be carried. + if _pyir_typed_update_is_advance(cur_value, [snap_value], region_op): + _pyir_raise_type_changed_in_region( + attr_name, snap_value.type, cur_value.type + ) + log().info( + "[pyir %s] leaf '%s' rebound to a differently-typed value " + "(%s -> %s); %s, not carried (%s)", + kind, + attr_name, + snap_value.type, + cur_value.type, + phrase_mismatch, + context, + ) + continue + # Leaves carry uniformly (all types); refuse only an update rooted at + # an IN-REGION allocation (entry-hoisted: aliases one buffer) -- the + # type-uniform escape judgment, same as the store choke. An admitted + # scratch handle (raw or re-loaded from the admitted slot) is neither + # refused nor re-carried: the slot cell already promotes into the + # carried phi, and the admission sweep owns its enforcement. + if _pyir_scratch_value_admitted(cur_value, region_op): + continue + if _pyir_memref_alloc_rooted_inside_region(cur_value, region_op): + _pyir_raise_memref_inregion_alloc_rebind(arg_name or attr_name) + try: + # 1. Ref acquisition (the ``scf.if`` caller shares one ref across + # both arms via *ref_cache*). + ref = ref_cache.get(id(snap_value)) if ref_cache is not None else None + if ref is None: + place = _pyir_region_carry_place_for(rebind_holder, attr_name, arg_name) + ref, _minted = _pyir_adopt_region_carry_ref( + place, snap_value, region_op + ) + if ref is None: + continue + if ref_cache is not None: + ref_cache[id(snap_value)] = ref + # 2. Region-entry load. A ``pyir.load`` of an opaque type comes back WRAPPED in its value-tree + # class; normalise to the backing SSA for the RAUW (a bare leaf is a no-op). + with ir.InsertionPoint.at_block_begin(region_block): + loaded = _raw_backing_ir_value(pyir.load(ref)) + if loaded is None: + continue + # 3. Region-scoped rewrite of in-region uses of the pre-region leaf to the loaded value, + # except the ``pyir.ref`` op (keep its init) and the region op itself (its OWN operands -- + # bounds/step/inits -- evaluate in the outer scope: rewriting a shared LB/UB SSA onto the + # load fails dominance). + ref_op = _get_defining_operation(ref) + load_op = _get_defining_operation(loaded) + region_op_norm = getattr(region_op, "operation", region_op) + exceptions = [ref_op] + for use in snap_value.uses: + user = getattr(use, "owner", None) + if user is None: + continue + user_op = getattr(user, "operation", user) + if user_op == load_op: + continue # the load itself was created with snap_value + if user_op == region_op_norm: + exceptions.append(user_op) + continue + if except_users_outside_region and not _op_is_inside_op( + user_op, region_op + ): + exceptions.append(user_op) + continue + # Except an in-region use the region-entry load does NOT dominate. + if not _pyir_load_dominates_use(loaded, user_op): + exceptions.append(user_op) + snap_value.replace_all_uses_except(loaded, exceptions) + # 4. Region-exit store of the mutated leaf, emitted at the CURRENT + # IP (the region block end, traced, before the terminator). + _pyir_emit_store(cur_value, ref) + except Exception as exc: + # Fail-loud wall: an un-carried mutated leaf silently pins every + # post-region read to the trapped in-region value. + raise DSLRuntimeError( + f"PyIR emission self-check: {phrase_wall} failed for " + f"value-tree leaf '{attr_name}' ({context}): {exc}" + ) from exc + # Record so the caller rebinds the carried value at the OUTER post-region IP (a + # load emitted here would land inside the still-open region). A + # whole-object RECONSTRUCT record always keeps the caller's slot + # index: an IMMUTABLE rebuilt compound (tuple/list) republishes into + # the ``mix_iter_args`` slot ONLY through it (in-place copy is + # impossible), and ``reconstruct_obj`` is always the top-level + # binding, so the index names the right slot. + if ( + self_leaf_arg_index_only + and reconstruct_obj is None + and not _is_self_leaf_record(rebind_holder, attr_name) + ): + rec_index = -1 + else: + rec_index = arg_index + records.append( + (rebind_holder, attr_name, ref, reconstruct_obj, leaf_index, rec_index) + ) + log().info( + "[pyir %s] carried updated value-tree leaf '%s' as %s (%s)", + kind, + attr_name, + phrase_carried, + context, + ) + + # A STABLE memref leaf written in place needs NO record: its carry is + # the backing memory, ordered by declared memory effects. + return records + + +def _pyir_carry_loop_body_mutated_leaves( + arg: Any, + snapshot: "list[tuple[Any, str, ir.Value]] | None", + loop_op: "ir.Operation", + body_block: "ir.Block", + context: str, + *, + opaque_only: bool = False, + arg_index: int = -1, + arg_name: "str | None" = None, +) -> "list[tuple[Any, str, ir.Value, Any, int, int]]": + """Retroactively carry a value-tree compound's UPDATED raw ``ir.Value`` + leaves through a ``pyir.ref`` so the C++ pass lifts each to a loop + iter_arg. Loop policy over the shared core: scalar leaves carry too + unless *opaque_only* (``scf.while``, which also skips leaves a nested + loop already carried); users outside the loop op are excepted from the + RAUW; only a SELF leaf or a whole-object RECONSTRUCT record keeps the + caller's *arg_index* (an in-place attribute record never needs it).""" + return _pyir_carry_region_mutated_leaves( + arg, + snapshot, + loop_op, + body_block, + context, + kind="loop", + opaque_only=opaque_only, + skip_inner_loop_carried=opaque_only, + except_users_outside_region=True, + self_leaf_arg_index_only=True, + arg_index=arg_index, + arg_name=arg_name, + ) + + +def _pyir_rebind_carried_leaves_post_loop( + records: "list[tuple]", +) -> "dict[int, Any]": + """Rebind each carried leaf to a dominating post-loop ``pyir.load %ref``.""" + if pyir is None: + return {} + # Group whole-object-rebind records by the object to reconstruct, tracking + # each object's ``mix_iter_args`` slot (the last non-negative ``arg_index``). + rebuilds: "dict[int, tuple[Any, dict[int, ir.Value], int]]" = {} + # Self-leaf rebinds: ``{arg_index: loaded_dsl_value}`` -- slot re-pointed at the WRAPPED load. + self_leaf_binds: "dict[int, Any]" = {} + for record in records: + holder, attr_name, ref, reconstruct_obj, leaf_index = record[:5] + arg_index = record[5] if len(record) > 5 else -1 + # SELF-LEAF: keep the WRAPPED ``pyir.load`` as the new slot binding so a post-region DSL read + # sees a proper typed value, not a bare SSA. + if _is_self_leaf_record(holder, attr_name): + try: + loaded_wrapped = pyir.load(ref) + except Exception as exc: + # Fail-loud wall: a swallowed post-loop load failure would + # silently leave the slot bound to the trapped in-loop value. + raise DSLRuntimeError( + "PyIR emission self-check: post-loop reload failed for " + f"the self-leaf slot '{attr_name}': {exc}" + ) from exc + if loaded_wrapped is not None and arg_index >= 0: + self_leaf_binds[arg_index] = loaded_wrapped + continue + try: + # ``pyir.load`` of an opaque type comes back WRAPPED (``_Tensor`` / + # ``_Pointer``); normalise to the backing SSA before rebinding. + loaded = _as_arith_capable_scalar_leaf( + _raw_backing_ir_value(pyir.load(ref)) + ) + if loaded is None: + continue + if reconstruct_obj is None: + _pyir_setattr_raw(holder, attr_name, loaded) + else: + key = id(reconstruct_obj) + _obj, leaves, ai = rebuilds.setdefault( + key, (reconstruct_obj, {}, arg_index) + ) + leaves[leaf_index] = loaded + if arg_index >= 0 and ai < 0: + rebuilds[key] = (_obj, leaves, arg_index) + except Exception as exc: + # Fail-loud wall: a swallowed store-back rebind failure would + # silently drop this leaf's carry (post-loop reads go stale). + raise DSLRuntimeError( + "PyIR emission self-check: post-loop leaf rebind failed for " + f"attribute '{attr_name}': {exc}" + ) from exc + # Reconstruct each whole-object-rebound compound from its loaded canonical + # leaves and copy the rebuilt state in place. + arg_index_to_obj: "dict[int, Any]" = {} + for obj, leaves, arg_index in rebuilds.values(): + try: + cur_leaves = _pyir_extract_any_leaf_values(obj) + if cur_leaves is None: + continue + # Substitute the carried leaves at their extract positions; any leaf + # the loop did not change keeps its current value. + new_vals = list(cur_leaves) + for li, lv in leaves.items(): + if 0 <= li < len(new_vals): + new_vals[li] = lv + rebuilt = _pyir_new_from_mlir_values_any(obj, new_vals) + if rebuilt is None: + continue + if _pyir_is_carryable_tuple(obj): + # A tuple/list is IMMUTABLE -- it cannot be copied in place. + if arg_index >= 0: + arg_index_to_obj[arg_index] = rebuilt + continue + # Copy the rebuilt instance state into ``obj`` so every alias observes + # the re-derived (carried) leaves. + for k, v in list(vars(rebuilt).items()): + _pyir_setattr_raw(obj, k, v) + if arg_index >= 0: + arg_index_to_obj[arg_index] = obj + except Exception as exc: + # Fail-loud wall: a swallowed reconstruct failure would silently + # leave every alias of the compound on its pre-loop leaves. + raise DSLRuntimeError( + "PyIR emission self-check: post-loop compound reconstruct " + f"failed for a {type(obj).__name__} instance: {exc}" + ) from exc + # Self-leaf rebinds re-point their slots too (the loaded wrapped value IS the new binding). + arg_index_to_obj.update(self_leaf_binds) + return arg_index_to_obj + + +def _pyir_rebind_local_place_cells_post_region( + mix_iter_args: "list", + mix_iter_arg_names: "list[str]", + op: "ir.Operation", +) -> None: + """Post-region rebind FROM the ledger cell (clause A, read half).""" + if pyir is None: + return + for idx in range(min(len(mix_iter_args), len(mix_iter_arg_names))): + try: + place = _pyir_local_place_for_name(mix_iter_arg_names[idx]) + if place is None: + continue + ref = _slot_refs.get(place) + if ref is None: + continue + try: + pointee = ref.type.pointee + except Exception: + continue + # OPAQUE pointees only -- scalar places are owned by D1/M2S. + if not _pyir_ir_type_is_opaque(pointee): + continue + if not _value_dominates_current_ip(ref): + continue + if not _pyir_slot_stored_in_body(ref, op): + continue + # Already the cell's dominating load -> nothing to do. + cur_raw = _raw_backing_ir_value(mix_iter_args[idx]) + if isinstance(cur_raw, ir.Value): + cur_owner = getattr(cur_raw, "owner", None) + if cur_owner is not None and not isinstance(cur_owner, ir.Block): + cur_op = getattr(cur_owner, "operation", cur_owner) + if ( + str(getattr(cur_op, "name", "")) == "pyir.load" + and _same_ir_value(cur_op.operands[0], ref) + and _value_dominates_current_ip(cur_raw) + ): + continue + loaded = pyir.load(ref) + if loaded is None: + continue + mix_iter_args[idx] = loaded + log().info( + "[pyir %s] post-region rebind of local '%s' from its place cell", + getattr(getattr(op, "operation", op), "name", "region"), + mix_iter_arg_names[idx], + ) + except Exception: + continue + + +def _pyir_rebind_unstored_scalar_carries_post_loop( + records: "list[tuple[Any, str, MutableValue]]", +) -> None: + """Rebind each stored-back scalar carry to a dominating post-loop ``pyir.load``.""" + if pyir is None: + return + for holder, attr_name, mv in records: + try: + # Row-authoritative reload: reconstruct from the store-time template. + loaded = mv.load() + _attach_mutable_ref(loaded, mv, f"post-loop carry rebind '{attr_name}'") + _pyir_setattr_raw(holder, attr_name, loaded) + except Exception: + continue + + +def _pyir_store_back_unstored_slot_carries_sweep( + loop_op: "ir.Operation", + body_block: "ir.Block", + context: str, +) -> "list[tuple[Any, str, MutableValue]]": + """Region-close store-back driven by the central slot-holder registry.""" + records: "list[tuple[Any, str, MutableValue]]" = [] + if pyir is None: + return records + # Id-keyed registry entries are a ``weakref.ref`` or the holder itself + # (non-weakrefable ``__slots__`` class held strongly for the trace). + for entry in list(_PYIR_SLOT_HOLDERS.values()): + holder = entry() if isinstance(entry, _weakref.ref) else entry + if holder is None: + continue + for slot_name, mv in _iter_owner_slot_mvs(holder): + # Direct attribute slots only: a composite key (``_PlaceSeg``) is + # a value-tree element, not a ``getattr``-able attribute. + if not isinstance(slot_name, str): + continue + try: + cur_value = getattr(holder, slot_name) + except AttributeError: + continue + # Only a numeric scalar leaf advances via the un-instrumented shape; an + # opaque leaf carries through _pyir_carry_loop_body_mutated_leaves. + if not _is_numeric_leaf_holder(cur_value): + continue + # A value carrying a load version is a fresh ``pyir.load`` refresh of the + # slot (set by MutableValue.load), not a mutation -- storing it back would + if getattr(cur_value, _PYIR_LOAD_VERSION_ATTR, None) is not None: + continue + # The row's own published representative, still unchanged in + # storage, is a declared non-mutation: the cell was minted from + # this very value at the publish choke, so there is no + # un-instrumented advance to store back (and its raw may be + # region-interior -- a body-exit store would serve SSA outside + # its defining region). + if getattr(cur_value, "_pyir_publish_representative_of", None) is mv: + continue + # The advanced value must be a fresh SSA defined inside the loop body. A + # read-only carried object keeps its pre-loop SSA (defined outside). + cur_raw = _raw_backing_ir_value(cur_value) + if cur_raw is None or not _ir_value_defined_inside_op(cur_raw, loop_op): + continue + # Unify onto the EXISTING slot ref the body loads from -- never mint a + # competing ref. + if mv is None or not mv._is_ref_accessible(): + continue + ref = mv.ref + if not isinstance(ref, ir.Value): + continue + if _ir_value_defined_inside_op(ref, loop_op): + continue # an in-body lazily-created ref cannot anchor a pre-loop carry + if not _pyir_ref_loaded_inside_op(ref, loop_op): + continue # never read in-body -> not a genuine carry, leave it alone + if _pyir_slot_stored_in_body(ref, loop_op): + continue # instrumented mutator already stores -> no-op + if _pyir_ref_pointee_type_changed(ref, cur_value): + # An scf iter_arg is type-invariant: a slot rebound to a DIFFERENTLY- + # TYPED value cannot be stored back. + _ref_loads = [] + try: + for _use in ref.uses: + _user = getattr(_use, "owner", None) + if _user is None: + continue + _user_op = getattr(_user, "operation", _user) + if getattr( + _user_op, "name", None + ) == "pyir.load" and _op_is_inside_op(_user_op, loop_op): + _ref_loads.extend(list(_user_op.results)) + except Exception: + # An unwalkable use list is UNKNOWN: refuse loudly rather + # than silently skip the carry. + _pyir_raise_type_changed_in_region(slot_name, None, cur_raw.type) + if _pyir_typed_update_is_advance(cur_raw, _ref_loads, loop_op): + _old_t = None + try: + _old_t = ref.type.pointee + except Exception: + pass + _pyir_raise_type_changed_in_region(slot_name, _old_t, cur_raw.type) + continue + # Wrapper-half leg comparison (V-4): over ONE MLIR type, an advance + # that changed the wrapper class has no Python join type -- a + # post-region read cannot reconstruct one consistent value. + if not _types_match(mv._value, cur_value) and not ( + isinstance(cur_value, type(mv._value)) + or isinstance(mv._value, type(cur_value)) + ): + raise DSLUserCodeError( + DiagId.WRAPPER_CLASS_MERGE, + var=str(slot_name), + old_cls=type(mv._value).__name__, + new_cls=type(cur_value).__name__, + ) + try: + # Body-exit store of the advanced scalar (the iter_arg yield); + # at the body-block end the store dominates the yield point. + _pyir_emit_store(cur_raw, ref) + except Exception as exc: + log().info( + "[pyir loop] registry-sweep slot store-back emit failed: %s", exc + ) + continue + records.append((holder, slot_name, mv)) + log().info( + "[pyir loop] stored un-instrumented numeric-slot advance '%s' back " + "into its existing slot ref as a loop iter_arg via registry sweep (%s)", + slot_name, + context, + ) + return records + + +def _pyir_carry_if_region_mutated_leaves( + arg: Any, + snapshot: "list[tuple[Any, str, ir.Value]] | None", + if_op: "ir.Operation", + region_block: "ir.Block", + ref_cache: "dict[int, ir.Value]", + context: str, + arg_index: int = -1, + exclude_aliased_leaves: "list[ir.Value] | None" = None, + arg_name: "str | None" = None, +) -> "list[tuple[Any, str, ir.Value, Any, int, int]]": + """Carry an ``scf.if`` branch's in-place-MUTATED value-tree leaf forward + as an ``scf.if`` result -- the read-side companion of the loop carry. + If policy over the shared core: OPAQUE leaves only (a scalar leaf is + already carried by the M2S function-entry slot); leaves aliased by a + captured free variable are excluded (*exclude_aliased_leaves*); both + arms share one ref per pre-``if`` SSA via *ref_cache*; every record + keeps the caller's *arg_index*.""" + return _pyir_carry_region_mutated_leaves( + arg, + snapshot, + if_op, + region_block, + context, + kind="if", + opaque_only=True, + skip_inner_loop_carried=False, + except_users_outside_region=False, + self_leaf_arg_index_only=False, + ref_cache=ref_cache, + exclude_aliased_leaves=exclude_aliased_leaves, + arg_index=arg_index, + arg_name=arg_name, + ) + + +def _pyir_snapshot_region_meta(arg: Any) -> "list[tuple[Any, str, Any]] | None": + """Snapshot *arg*'s meta-scalar leaf holders before a region body.""" + holders = _pyir_walk_meta_holders(arg) + return holders if holders else None + + +def _pyir_snapshot_meta_numeric_leaves( + arg: Any, +) -> "list[tuple[Any, str, Any]] | None": + """Snapshot *arg*'s parent-level META Numeric leaf holders before a region. Consumed + by :func:`_pyir_restore_meta_numeric_leaves`.""" + holders = _walk_meta_numeric_leaf_holders(arg) + return holders if holders else None + + +def _pyir_loop_carried_meta_needs_rebind( + value: Any, ref: "ir.Value", loop_op: "ir.Operation" +) -> bool: + """Return True for a loop-carried META scalar whose slot the body STORED into.""" + if pyir is None: + return False + try: + if not isinstance(value, _WatchedM): + return False + if not isinstance(ref, ir.Value): + return False + return _pyir_slot_stored_in_body(ref, loop_op) + except Exception: + return False + + +def _ensure_leaf_ref_no_load( + owner: Any, slot_name: Any, leaf: Any, context: str +) -> Any: + """Ensure a staged leaf has a ``pyir.ref`` slot under ``(owner, slot_name)`` + without emitting a load (the single load is deferred to the consumer).""" + if not (_is_staged_value(leaf) and _can_carry_leaf_ref(leaf)): + return leaf + existing = _get_slot_mv(owner, slot_name) + if existing is not None and existing._is_ref_accessible(): + return leaf # already carried via a registered slot + if _pyir_value_tracked_by_accessible_ref(leaf): + return leaf # already carried via the leaf's _mutable_ref + try: + leaf = _fresh_wrapper(leaf) + mv = _create_ref(leaf) + mv = _set_slot_mv(owner, slot_name, mv) + _attach_mutable_ref(leaf, mv, f"while-carried leaf '{context}'") + except Exception as exc: + # Fail-loud wall: a swallowed slot-creation failure would silently + # drop this leaf's carry (reads bake the stale pre-loop value). + raise DSLRuntimeError( + "PyIR emission self-check: slot creation failed for the " + f"while-carried leaf '{context}' (slot {slot_name!r}): {exc}" + ) from exc + return leaf + + +def _promote_tuple_leaf_refs( + owner: Any, + slot_name: Any, + t: tuple, + context: str, +) -> tuple: + """Ensure the staged leaves of a tuple have ``pyir.ref`` slots (no load), under + ``_decompose_tuple``'s slot keys; rebuilds through the tuple's own subtype.""" + result: list = [] + for i, elem in enumerate(t): + elem_slot = _place_seg_child(slot_name, i) if slot_name is not None else None + if _is_staged_value(elem) and _can_carry_leaf_ref(elem): + result.append( + _ensure_leaf_ref_no_load(owner, elem_slot, elem, f"{context}[{i}]") + ) + elif isinstance(elem, tuple): + result.append( + _promote_tuple_leaf_refs(owner, elem_slot, elem, f"{context}[{i}]") + ) + else: + result.append(elem) + return _rebuild_tuple_like(t, result) + + +def _promote_compound_leaf_refs( + obj: object, context: str, _visited: "set[int] | None" = None +) -> bool: + """Materialise ``pyir.ref`` slots at the FIRST READ of a loop-carried value-tree + compound so each staged leaf carries as an iter_arg (mutates *obj* in place).""" + if pyir is None or not is_inside_staged_cf(): + return False + # Break cycles in the value-tree object graph (a nested field reaching back + # to an ancestor). Re-promoting an already-visited object is a no-op. + if _visited is None: + _visited = set() + if id(obj) in _visited: + return False + _visited.add(id(obj)) + if not _implements_dynamic_expression(obj) or not _has_instance_storage(obj): + return False + if not _pyir_plain_storage_setattr(type(obj)): + return False + # Gate on "has any staged content", not "all fields decomposable": non-carryable + # fields are left for the value-tree reconstruct, carryable leaves still promote. + if not _has_any_staged_content(obj): + return False + + promoted_any = False + for attr_name in _get_instance_attrs(obj): + value = getattr(obj, attr_name) + if _is_staged_value(value) and _can_carry_leaf_ref(value): + # Ensure the leaf has a ``pyir.ref`` slot without a load, keyed on + # ``(obj, attr_name)`` so the bottom-of-body decompose reuses it. + fresh = _ensure_leaf_ref_no_load( + obj, attr_name, value, f"{context}.{attr_name}" + ) + if fresh is not value: + try: + _pyir_setattr_raw(obj, attr_name, fresh) + promoted_any = True + except (AttributeError, TypeError): + pass + elif isinstance(value, tuple): + # Only touch tuples carrying a staged leaf; a pure meta tuple + # (shapes/strides) must stay constant. + if any( + _is_staged_value(e) and _can_carry_leaf_ref(e) + for e in _flatten_tuple(value) + ): + new_tuple = _promote_tuple_leaf_refs( + obj, attr_name, value, f"{context}.{attr_name}" + ) + try: + _pyir_setattr_raw(obj, attr_name, new_tuple) + promoted_any = True + except (AttributeError, TypeError): + pass + elif ( + _has_instance_storage(value) + and not _is_staged_value(value) + and not isinstance(value, (int, float, bool, str, bytes, type)) + and _implements_dynamic_expression(value) + and _has_decomposable_staged_fields(value) + ): + if _promote_compound_leaf_refs(value, f"{context}.{attr_name}", _visited): + promoted_any = True + if promoted_any: + log().info("[pyir] promoted value-tree leaves for %s", context) + return promoted_any + + +# Static export surface (regenerate with scripts/gen_pyir_all.py): the +# module's OWN names -- its namespace minus what lower chain modules +# already export and minus the generator's INTERNAL_ONLY registry. +__all__ = [ + "_pyir_push_loop_body_scope", + "_pyir_pop_loop_body_scope", + "_meta_promote_slot", + "_stage_compound_value", + "_stage_meta_compound_leaves", + "_pyir_snapshot_region_arg", + "_pyir_carry_loop_body_mutated_leaves", + "_pyir_rebind_carried_leaves_post_loop", + "_pyir_rebind_local_place_cells_post_region", + "_pyir_rebind_unstored_scalar_carries_post_loop", + "_pyir_store_back_unstored_slot_carries_sweep", + "_pyir_carry_if_region_mutated_leaves", + "_pyir_snapshot_region_meta", + "_pyir_snapshot_meta_numeric_leaves", + "_pyir_loop_carried_meta_needs_rebind", + "_ensure_leaf_ref_no_load", + "_promote_compound_leaf_refs", +] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_preprocessor.py b/python/CuTeDSL/cutlass/base_dsl/pyir_preprocessor.py index 10ceb8b730..efbad6b987 100644 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_preprocessor.py +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_preprocessor.py @@ -18,20 +18,62 @@ import ast import contextlib +import types from collections.abc import Generator -from copy import deepcopy from dataclasses import dataclass, field +from typing import Any, Callable from typing_extensions import override from .ast_preprocessor import ( DSLPreprocessor, + OrderedSet, + Region, ScopeManager, _create_module_attribute, + _ComprehensionT, _deepcopy_ast_root, SessionData, ) -from .common import DSLUserCodeError +from .common import DSLRuntimeError, DSLUserCodeError from .diagnostics import DiagId +from .pyir_class_facts import ( + on_pyir_preprocess_session_start, +) + +# Imported for its module-attribute side effect too: the rewritten-function +# preamble binds ``__pyir_runtime__ = __base_dsl__.pyir_runtime``, so the +# submodule must be loaded before any rewritten function runs. +from . import pyir_runtime + +# Runtime-module alias bound once in every rewritten function's preamble: +# generated code names pyir_runtime symbols through their owning module. +_PYIR_RUNTIME_ALIAS = "__pyir_runtime__" + +# Emission-time fact: a symbol referenced through the runtime alias must be +# one the pyir_runtime facade exports (its non-dunder module namespace: the +# facade re-exports the chain surface and carries no ``__all__`` of its own). +_PYIR_RUNTIME_EXPORTS = frozenset( + n for n in vars(pyir_runtime) if not n.startswith("__") +) + + +def _create_runtime_attribute( + func_name: str, + *, + lineno: "int | None" = None, + col_offset: "int | None" = None, +) -> ast.Attribute: + """``__pyir_runtime__.`` with optional location info.""" + assert func_name in _PYIR_RUNTIME_EXPORTS, func_name + base = ast.Name(id=_PYIR_RUNTIME_ALIAS, ctx=ast.Load()) + result = ast.Attribute(value=base, attr=func_name, ctx=ast.Load()) + if lineno is not None and col_offset is not None: + for node in (base, result): + node.lineno = lineno + node.end_lineno = lineno + node.col_offset = col_offset + node.end_col_offset = col_offset + return result def _unparse_safe(node: ast.AST) -> str: @@ -44,6 +86,106 @@ def _unparse_safe(node: ast.AST) -> str: return type(node).__name__ +def _whole_name_rebound_names(body: "list[ast.stmt]") -> set[str]: + """Bare names the statements in *body* REBIND as a whole (assign / augassign / + for-target), lambda/nested-def bodies excluded; the list-carry promotion fact.""" + names: set[str] = set() + + def _record(t: "ast.expr") -> None: + stack: "list[ast.expr]" = [t] + while stack: + n = stack.pop() + if isinstance(n, (ast.Tuple, ast.List)): + stack.extend(n.elts) + elif isinstance(n, ast.Starred): + stack.append(n.value) + elif isinstance(n, ast.Name): + names.add(n.id) + + stack: "list[ast.AST]" = list(body) + while stack: + node = stack.pop() + if isinstance(node, (ast.Lambda, ast.FunctionDef, ast.AsyncFunctionDef)): + continue + if isinstance(node, ast.Assign): + for t in node.targets: + _record(t) + elif isinstance(node, (ast.AugAssign, ast.AnnAssign)): + _record(node.target) + elif isinstance(node, (ast.For, ast.AsyncFor)): + _record(node.target) + stack.extend(ast.iter_child_nodes(node)) + return names + + +def _locally_bound_names(node: ast.AST) -> set[str]: + """Names bound only in a nested comprehension/lambda scope inside *node*; + a statement-scope hoist must never reference them.""" + bound: set[str] = set() + for child in ast.walk(node): + if isinstance( + child, (ast.ListComp, ast.SetComp, ast.DictComp, ast.GeneratorExp) + ): + for gen in child.generators: + for tgt in ast.walk(gen.target): + if isinstance(tgt, ast.Name) and isinstance(tgt.ctx, ast.Store): + bound.add(tgt.id) + elif isinstance(child, ast.Lambda): + args = child.args + for arg in ( + *args.posonlyargs, + *args.args, + *args.kwonlyargs, + ): + bound.add(arg.arg) + if args.vararg is not None: + bound.add(args.vararg.arg) + if args.kwarg is not None: + bound.add(args.kwarg.arg) + return bound + + +# Emission budget for one access chain's hop reads: chains are instrumented at +# their exact AST depth; a deeper chain refuses loudly instead of truncating. +_READ_CHAIN_HOP_BUDGET = 64 + + +def _ast_node_depths(node: ast.AST) -> dict[int, int]: + """Depth of every AST node id under *node*: a chain nested inside another + chain's base (call argument, subscript key) has strictly greater depth.""" + depths: dict[int, int] = {} + stack: list[tuple[ast.AST, int]] = [(node, 0)] + while stack: + cur, d = stack.pop() + depths[id(cur)] = d + for child in ast.iter_child_nodes(cur): + stack.append((child, d + 1)) + return depths + + +def _deferred_execution_node_ids(node: ast.AST) -> set[int]: + """Ids of AST nodes whose evaluation Python defers past the statement: + lambda bodies and a genexp's lazy parts (elt / targets / ifs / inner iters).""" + deferred: set[int] = set() + for child in ast.walk(node): + if isinstance(child, ast.Lambda): + for sub in ast.walk(child.body): + deferred.add(id(sub)) + elif isinstance(child, ast.GeneratorExp): + for sub in ast.walk(child.elt): + deferred.add(id(sub)) + for pos, gen in enumerate(child.generators): + for sub in ast.walk(gen.target): + deferred.add(id(sub)) + if pos > 0: + for sub in ast.walk(gen.iter): + deferred.add(id(sub)) + for cond in gen.ifs: + for sub in ast.walk(cond): + deferred.add(id(sub)) + return deferred + + @dataclass class PyIRScopeManager(ScopeManager): """ @@ -101,6 +243,410 @@ class PyIRSessionData(SessionData): """ scope_manager: PyIRScopeManager = field(default_factory=PyIRScopeManager.create) + # Synthetic live-out carry name per loop node, keyed by ``id(node)`` so + # nested loops don't collide; forces a body-entry promotion of the carry. + liveout_loop_carried_vars: dict[int, str] = field(default_factory=dict) + + +def _collect_direct_attr_write_pairs(body: "list[ast.stmt]") -> "set[tuple[str, str]]": + """``(base, attr)`` pairs the statements in *body* assign directly with a + Name-rooted base; a nested target (``c.sub.n``) records the dotted base + (``("c.sub", "n")``). Lambda/nested-def bodies excluded (they bind at + call time).""" + pairs: "set[tuple[str, str]]" = set() + + def _record(t: "ast.expr") -> None: + stack: "list[ast.expr]" = [t] + while stack: + n = stack.pop() + if isinstance(n, (ast.Tuple, ast.List)): + stack.extend(n.elts) + elif isinstance(n, ast.Starred): + stack.append(n.value) + elif isinstance(n, ast.Attribute): + base = _name_rooted_receiver_path(n.value) + if base is not None: + pairs.add((base, n.attr)) + + stack: "list[ast.AST]" = list(body) + while stack: + node = stack.pop() + if isinstance(node, (ast.Lambda, ast.FunctionDef, ast.AsyncFunctionDef)): + continue + if isinstance(node, ast.Assign): + for t in node.targets: + _record(t) + elif isinstance(node, (ast.AugAssign, ast.AnnAssign)): + _record(node.target) + elif isinstance(node, (ast.For, ast.AsyncFor)): + _record(node.target) + stack.extend(ast.iter_child_nodes(node)) + return pairs + + +def _name_rooted_receiver_path(node: "ast.expr") -> "str | None": + """Dotted receiver path when *node* is a Name-rooted attribute chain + (``a`` or ``a.b.c``); ``None`` for any other receiver shape.""" + parts: "list[str]" = [] + while isinstance(node, ast.Attribute): + parts.append(node.attr) + node = node.value + if not isinstance(node, ast.Name): + return None + parts.append(node.id) + return ".".join(reversed(parts)) + + +def _collect_receiver_method_call_pairs( + body: "list[ast.stmt]", +) -> "set[tuple[str, str]]": + """``(base, method)`` pairs the statements in *body* call with a Name-rooted + receiver path (``base.method(...)``, ``root.attr.method(...)``); the base is + the dotted path. Lambda/nested-def bodies excluded (they bind at call + time). The callee's receiver-attr write facts complete each pair into + region write facts at region entry.""" + pairs: "set[tuple[str, str]]" = set() + stack: "list[ast.AST]" = list(body) + while stack: + node = stack.pop() + if isinstance(node, (ast.Lambda, ast.FunctionDef, ast.AsyncFunctionDef)): + continue + if isinstance(node, ast.Call): + f = node.func + if isinstance(f, ast.Attribute): + base = _name_rooted_receiver_path(f.value) + if base is not None: + pairs.add((base, f.attr)) + stack.extend(ast.iter_child_nodes(node)) + return pairs + + +def _collect_free_call_arg_pairs( + body: "list[ast.stmt]", +) -> "set[tuple[str, str]]": + """``(func_name, arg_base)`` pairs the statements in *body* call as a + bare-Name free function with a Name-rooted first argument (``step(c)``); + lambda/nested-def bodies excluded (they bind at call time). The callee's + first-parameter write facts complete each pair into region write facts at + region entry.""" + pairs: "set[tuple[str, str]]" = set() + stack: "list[ast.AST]" = list(body) + while stack: + node = stack.pop() + if isinstance(node, (ast.Lambda, ast.FunctionDef, ast.AsyncFunctionDef)): + continue + if isinstance(node, ast.Call): + f = node.func + if isinstance(f, ast.Name) and node.args: + base = _name_rooted_receiver_path(node.args[0]) + if base is not None: + pairs.add((f.id, base)) + stack.extend(ast.iter_child_nodes(node)) + return pairs + + +# F-COVER: version of the choke set this preprocessor emits. Stamped as +# ``__pyir_rewritten__`` on every rewritten function object; bump on any +# change to the emitted choke vocabulary so stale rewrites fail attestation. +PYIR_CHOKE_SET_VERSION: int = 6 + + +def _own_scope_nonlocal_names(body: "list[ast.stmt]") -> "list[str]": + """Names declared ``nonlocal`` by THIS function's own statements, in + first-seen order; nested function/class scopes declare for themselves.""" + names: "list[str]" = [] + queue: "list[ast.AST]" = list(body) + i = 0 + while i < len(queue): + stmt = queue[i] + i += 1 + if isinstance( + stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.Lambda, ast.ClassDef) + ): + continue + if isinstance(stmt, ast.Nonlocal): + for n in stmt.names: + if n not in names: + names.append(n) + continue + queue.extend(ast.iter_child_nodes(stmt)) + return names + + +class _ScopeBindingCollector(ast.NodeVisitor): + """Collect the names ONE function scope binds anywhere in its body + (CPython symbol-table locality rule), without descending into nested + function/class/lambda scopes or comprehension scopes.""" + + def __init__(self) -> None: + self.bound: set[str] = set() + self.declared_global: set[str] = set() + self.declared_nonlocal: set[str] = set() + + def visit_Name(self, node: ast.Name) -> None: + if isinstance(node.ctx, (ast.Store, ast.Del)): + self.bound.add(node.id) + + def visit_FunctionDef(self, node: ast.FunctionDef) -> None: + self.bound.add(node.name) # the def binds its name; its body is a new scope + + def visit_AsyncFunctionDef(self, node: "ast.AsyncFunctionDef") -> None: + self.bound.add(node.name) + + def visit_ClassDef(self, node: ast.ClassDef) -> None: + self.bound.add(node.name) + + def visit_Lambda(self, node: ast.Lambda) -> None: + pass # inner scope + + def visit_ListComp(self, node: ast.ListComp) -> None: + pass # comprehension targets live in their own scope + + visit_SetComp = visit_ListComp # type: ignore[assignment] + visit_DictComp = visit_ListComp # type: ignore[assignment] + visit_GeneratorExp = visit_ListComp # type: ignore[assignment] + + def visit_Import(self, node: ast.Import) -> None: + for alias in node.names: + self.bound.add((alias.asname or alias.name).split(".")[0]) + + def visit_ImportFrom(self, node: ast.ImportFrom) -> None: + for alias in node.names: + self.bound.add(alias.asname or alias.name) + + def visit_ExceptHandler(self, node: ast.ExceptHandler) -> None: + if node.name: + self.bound.add(node.name) + self.generic_visit(node) + + def visit_MatchAs(self, node: "ast.MatchAs") -> None: + if node.name: + self.bound.add(node.name) + self.generic_visit(node) + + def visit_MatchStar(self, node: "ast.MatchStar") -> None: + if node.name: + self.bound.add(node.name) + self.generic_visit(node) + + def visit_Global(self, node: ast.Global) -> None: + self.declared_global.update(node.names) + + def visit_Nonlocal(self, node: ast.Nonlocal) -> None: + self.declared_nonlocal.update(node.names) + + +def _function_scope_bindings( + fn_node: "ast.FunctionDef | ast.AsyncFunctionDef", +) -> set[str]: + """All names *fn_node*'s scope owns anywhere in its body: parameters plus + collected bindings, minus ``global`` declarations; a name the scope itself + declares ``nonlocal`` resolves further up a chain CPython already + validated when the user's module compiled, so it also counts as owned.""" + collector = _ScopeBindingCollector() + for stmt in fn_node.body: + collector.visit(stmt) + args = fn_node.args + for arg in args.posonlyargs + args.args + args.kwonlyargs: + collector.bound.add(arg.arg) + if args.vararg is not None: + collector.bound.add(args.vararg.arg) + if args.kwarg is not None: + collector.bound.add(args.kwarg.arg) + return (collector.bound | collector.declared_nonlocal) - collector.declared_global + + +def _is_synthetic_call(func_node: ast.expr) -> bool: + """True for a direct preprocessor-emitted helper invocation (an + attribute chain rooted at ``__base_dsl__``/``__module_dsl__`` or the + runtime-module alias). A call WHOSE CALLEE is such a helper call is a + user call over a choke-produced value and stays wrapped; + machinery-built outer calls mark themselves ``_pyir_synth`` at + construction.""" + node = func_node + while isinstance(node, ast.Attribute): + node = node.value + return isinstance(node, ast.Name) and node.id in ( + "__base_dsl__", + "__module_dsl__", + _PYIR_RUNTIME_ALIAS, + ) + + +class _CallBoundaryWrapPass(ast.NodeTransformer): + """Post-instrumentation pass turning ``f(args)`` into + ``_pyir_call_boundary_(f)(args)`` (the pyir_runtime dispatcher, bound as + an alias in the rewritten-function preamble) after the main visitor, so + no pattern matcher sees the wrapper; synthetic calls skipped, and + conditioned short-circuit operands carry the ``sc_rhs`` fact.""" + + _SHORT_CIRCUIT_HELPERS = ("and_", "or_") + + # Bare-name builtin calls the MAIN VISITOR emits in synthetic statements + # (marked ``_pyir_synth`` where they carry user expressions) plus builtins + # whose behavior depends on the calling frame (``super``) or that are pure + # w.r.t. tracked places, so wrapping them only adds dispatch cost. + _BARE_BUILTIN_SKIP = frozenset({"range", "type", "super", "locals"}) + + # Zero-arg frame-reflection builtins: rewritten to the position-aware + # ``_pyir_traced_locals`` choke (the synthesized-scope fact is a rewrite- + # time fact, so it is passed as a constant at emission). + _FRAME_REFLECTION_BUILTINS = frozenset({"locals", "vars"}) + + # Dispatcher alias bound once at function entry (one local/closure load per + # wrapped call); a nested CLASS body uses the module-attribute form instead. + ALIAS_NAME = "_pyir_call_boundary_" + + def __init__(self) -> None: + super().__init__() + self._sc_depth = 0 + self._class_depth = 0 + # Scope-kind stack over nested defs/lambdas met in the body: True for + # a preprocessor-SYNTHESIZED scope (arm/body block defs), False for a + # user-authored nested scope (its locals() is its own frame -- truth). + self._scope_synth: list[bool] = [] + self.wrapped_any = False + # Identity memo: each node OBJECT is transformed exactly once. On a + # proper tree this never fires; if the input degenerates to a DAG a + # re-entered node returns its first result instead of being re-wrapped + # once per path (the memo value pins the key object, keeping ids + # stable for the pass lifetime). + self._memo: dict[int, tuple[ast.AST, ast.AST]] = {} + + def visit(self, node: ast.AST) -> Any: + entry = self._memo.get(id(node)) + if entry is not None and entry[0] is node: + return entry[1] + result = super().visit(node) + self._memo[id(node)] = (node, result) + return result + + def run_on_function_body(self, func_def: ast.FunctionDef) -> bool: + """Wrap calls in *func_def*'s BODY only (decorators/defaults evaluate + outside the alias's scope); True when at least one call was wrapped.""" + func_def.body = [self.visit(stmt) for stmt in func_def.body] + return self.wrapped_any + + def visit_ClassDef(self, node: ast.ClassDef) -> ast.ClassDef: + self._class_depth += 1 + try: + self.generic_visit(node) + finally: + self._class_depth -= 1 + return node + + def _visit_scope(self, node: Any) -> Any: + """Visit a nested def/lambda: decorators and argument defaults evaluate + in the ENCLOSING scope; only the body runs in the new scope.""" + if hasattr(node, "decorator_list"): + node.decorator_list = [self.visit(d) for d in node.decorator_list] + args = node.args + args.defaults = [self.visit(d) for d in args.defaults] + args.kw_defaults = [ + self.visit(d) if d is not None else None for d in args.kw_defaults + ] + self._scope_synth.append(getattr(node, "_pyir_synth_scope", False)) + try: + if isinstance(node, ast.Lambda): + node.body = self.visit(node.body) + else: + node.body = [self.visit(stmt) for stmt in node.body] + finally: + self._scope_synth.pop() + return node + + def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef: + return self._visit_scope(node) + + def visit_AsyncFunctionDef( + self, node: ast.AsyncFunctionDef + ) -> ast.AsyncFunctionDef: + return self._visit_scope(node) + + def visit_Lambda(self, node: ast.Lambda) -> ast.Lambda: + return self._visit_scope(node) + + @staticmethod + def _short_circuit_helper_name(func: ast.expr) -> bool: + # ``__module_dsl__.and_`` / bare ``and_`` (module-attribute form + # emitted by visit_BoolOp). + if isinstance(func, ast.Attribute): + return func.attr in _CallBoundaryWrapPass._SHORT_CIRCUIT_HELPERS + if isinstance(func, ast.Name): + return func.id in _CallBoundaryWrapPass._SHORT_CIRCUIT_HELPERS + return False + + def _wrap_frame_reflection(self, node: ast.Call) -> ast.Call: + """``locals()`` / zero-arg ``vars()`` -> ``_pyir_traced_locals( + locals(), locals, )``: the mapping still evaluates in the + calling frame; the callee load lets the choke verify the builtin + (a shadowed name passes through); the synthesized-scope fact is a + rewrite-time constant.""" + node._pyir_synth = True # type: ignore[attr-defined] + callee_name = node.func.id # type: ignore[attr-defined] + callee_load = ast.copy_location(ast.Name(id=callee_name, ctx=ast.Load()), node) + in_synth = bool(self._scope_synth and self._scope_synth[-1]) + wrapper = ast.Call( + func=_create_module_attribute( + "_pyir_traced_locals", + submodule_name="pyir_runtime", + lineno=getattr(node, "lineno", None), + col_offset=getattr(node, "col_offset", None), + ), + args=[node, callee_load, ast.Constant(value=in_synth)], + keywords=[], + ) + return ast.copy_location(wrapper, node) + + def visit_Call(self, node: ast.Call) -> ast.Call: + if self._short_circuit_helper_name(node.func) and len(node.args) >= 2: + # First operand is unconditioned (Python evaluates it always); + # the remaining operands are the short-circuited RHS. + node.args[0] = self.visit(node.args[0]) + self._sc_depth += 1 + try: + node.args[1:] = [self.visit(a) for a in node.args[1:]] + node.keywords = [self.visit(kw) for kw in node.keywords] + finally: + self._sc_depth -= 1 + return node + self.generic_visit(node) + if getattr(node, "_pyir_synth", False): + return node + if _is_synthetic_call(node.func): + return node + if ( + isinstance(node.func, ast.Name) + and node.func.id in self._FRAME_REFLECTION_BUILTINS + and not node.args + and not node.keywords + ): + return self._wrap_frame_reflection(node) + if isinstance(node.func, ast.Name) and node.func.id in self._BARE_BUILTIN_SKIP: + return node + wrap_args: list[ast.expr] = [node.func] + if self._sc_depth > 0: + wrap_args.append(ast.Constant(value=True)) + if self._class_depth == 0: + wrap_func: ast.expr = ast.Name(id=self.ALIAS_NAME, ctx=ast.Load()) + ast.copy_location(wrap_func, node) + else: + wrap_func = _create_module_attribute( + "_pyir_call_boundary_", + submodule_name="pyir_runtime", + lineno=getattr(node, "lineno", None), + col_offset=getattr(node, "col_offset", None), + ) + node.func = ast.copy_location( + ast.Call( + func=wrap_func, + args=wrap_args, + keywords=[], + ), + node, + ) + self.wrapped_any = True + return node class PyIRDSLPreprocessor(DSLPreprocessor): @@ -114,6 +660,17 @@ class PyIRDSLPreprocessor(DSLPreprocessor): - Tuple unpacking decomposition """ + # F-COVER carrier: the rewrite point stamps this version on the function + # object; the base (non-PyIR) preprocessor advertises no choke set. + choke_set_version: "int | None" = PYIR_CHOKE_SET_VERSION + + @override + def _start_session(self) -> None: + # Materialize decoration-recorded per-class self-field write facts before + # any PyIR tracing, so the registry is complete when a region consults it. + on_pyir_preprocess_session_start() + super()._start_session() + @override def _create_session_data(self) -> SessionData: return PyIRSessionData() @@ -133,98 +690,267 @@ def _create_closure_check_call( # PYIRToSCF promotes to iter_args, so the check is unnecessary. return None - @override - def _create_lambda_check_call( - self, called_value_symbols: list[str], node: ast.stmt - ) -> ast.Expr | None: - # Emit ``lambda_capture_check([names...])`` at staged-region entry. - # A raw lambda invoked in the region reads its captures from the - # ENCLOSING frame, not the region's loop-carried slots (pyir's ref - # mechanism does not reach an un-preprocessed lambda body), so a - # captured meta mutated inside the region would freeze at its - # first-pass value — a silent miscompile. The check rejects that at - # trace time (non-lambdas/jit-lambdas/region-local lambdas pass). - # - # This override lives ONLY on the PyIR subclass (instantiated only when - # pyir is enabled, see preprocess_mode). The base returns None, so - # non-pyir compilation emits a byte-identical region with no - # lambda_capture_check call — pyir-gated by construction. - if not called_value_symbols: - return None - return ast.Expr( - ast.Call( + def _refuse_construct(self, diag: DiagId, node: ast.AST) -> None: + """Curated refusal for a construct the rewrite cannot compile.""" + raise DSLUserCodeError( + diag, + filename=self.session_data.file_name, + lineno=getattr(node, "lineno", None), + col_offset=getattr(node, "col_offset", None), + end_col_offset=getattr(node, "end_col_offset", None), + ) + + # A generator or coroutine body never runs at its call site (calling only + # mints the generator/coroutine object), so a rewritten def containing one + # of these constructs would compile to a body the trace can never observe. + # The wall is a PyIR-mode fact: the base rewrite compiles these constructs + # exactly as before. + def visit_Yield(self, node: ast.Yield) -> None: + self._refuse_construct(DiagId.UNSUP_YIELD, node) + + def visit_YieldFrom(self, node: ast.YieldFrom) -> None: + self._refuse_construct(DiagId.UNSUP_YIELD, node) + + def visit_AsyncFunctionDef(self, node: "ast.AsyncFunctionDef") -> None: + self._refuse_construct(DiagId.UNSUP_ASYNC, node) + + def visit_Await(self, node: "ast.Await") -> None: + self._refuse_construct(DiagId.UNSUP_ASYNC, node) + + def visit_AsyncFor(self, node: "ast.AsyncFor") -> None: + self._refuse_construct(DiagId.UNSUP_ASYNC, node) + + def visit_AsyncWith(self, node: "ast.AsyncWith") -> None: + self._refuse_construct(DiagId.UNSUP_ASYNC, node) + + def visit_TryStar(self, node: ast.AST) -> None: + self._refuse_construct(DiagId.UNSUP_EXCEPT_STAR, node) + + # Innermost-last stack of the function defs being visited: the whole-scope + # owner table for ``nonlocal`` name resolution (a PyIR-mode fact; the base + # rewrite resolves nonlocal from the visitation position only). + _function_scope_stack: "tuple[ast.FunctionDef, ...]" = () + + def visit_Nonlocal(self, node: ast.Nonlocal) -> ast.Nonlocal: + active_symbols = self.session_data.scope_manager.get_active_symbols() + nonlocal_names = OrderedSet(node.names) + intersect = nonlocal_names.intersections(active_symbols) + # Ownership is a whole-scope fact: an enclosing scope's binding owns + # the name even when it appears textually after this nested def. + enclosing_bindings: "list[set[str]] | None" = None + for name in node.names: + if name in intersect: + continue + if enclosing_bindings is None: + enclosing_bindings = [ + _function_scope_bindings(fn) + for fn in self._function_scope_stack[:-1] + ] + if any(name in bound for bound in enclosing_bindings): + continue + raise DSLUserCodeError( + DiagId.UNSUP_NONLOCAL, + filename=self.session_data.file_name, + lineno=getattr(node, "lineno", None), + col_offset=getattr(node, "col_offset", None), + end_col_offset=getattr(node, "end_col_offset", None), + stmt=ast.unparse(node), + name=name, + ) + self.generic_visit(node) + return node + + def _expand_boolop_evaluate_once(self, node: ast.BoolOp) -> ast.expr: + # Visit child nodes first + self.generic_visit(node) + + # Short-circuit evaluation expands explicitly, like the base rewrite, + # but in the evaluate-once form (a PyIR-mode fact: the call-boundary + # post-pass re-visits shared nodes, so the emitted module must stay a + # tree, and a watched meta bool must evaluate its producer once). + if isinstance(node.op, ast.And): + # Emitted form (one lhs evaluation, Python short-circuit): + # tmp if bool_short_circuits((tmp := lhs), False) else and_(tmp, rhs) + short_circuit_value = ast.Constant(value=False) + self.session_data.import_top_module = True + elif isinstance(node.op, ast.Or): + # Emitted form (one lhs evaluation, Python short-circuit): + # tmp if bool_short_circuits((tmp := lhs), True) else or_(tmp, rhs) + short_circuit_value = ast.Constant(value=True) + self.session_data.import_top_module = True + else: + # BoolOp should be either And or Or -- reaching here is a compiler + # bug (the AST grammar only produces And/Or), not an author mistake. + raise DSLRuntimeError( + f"Unsupported boolean operation: {node.op}", + filename=self.session_data.file_name, + snippet=ast.unparse(node), + ) + + # Evaluate-once lowering: the lhs binds to a synthesized temp INSIDE + # the test (a named expression), so every synthesized position reads + # that one evaluation -- the re-evaluating form runs the lhs per + # position and deepcopies it per chain link (super-linear on chains, + # and a side-effecting lhs fires more than once). The test asks + # ``bool_short_circuits`` whether the temp is a Python-truth bool (a + # plain bool or a watched meta bool) equal to the short-circuit + # value; only the non-short-circuit arm evaluates the rhs, exactly + # like Python. Every synthesized node is built fresh per position + # (the emitted module must stay a tree). + helper_name = "and_" if isinstance(node.op, ast.And) else "or_" + sc_value = short_circuit_value.value + + lhs = node.values[0] + for i in range(1, len(node.values)): + tmp_name = f"_pyir_bool_{self.session_data.counter}" + self.session_data.counter += 1 + test = ast.Call( func=_create_module_attribute( - "lambda_capture_check", + "bool_short_circuits", lineno=node.lineno, col_offset=node.col_offset, ), args=[ - ast.List( - elts=[ - ast.Name(id=c, ctx=ast.Load()) for c in called_value_symbols - ], - ctx=ast.Load(), - ) + ast.NamedExpr( + target=ast.Name(id=tmp_name, ctx=ast.Store()), + value=lhs, + ), + ast.Constant(value=sc_value), ], keywords=[], ) - ) - - def visit_With(self, node: ast.With) -> ast.AST: - # Base handling first: register optional-vars in scope + recurse into - # children. Then wrap each context-manager expression in the trace-time - # guard ``with_ctxmgr_check`` so a ``with`` whose __enter__/__exit__ are - # raw-Python user dunders is rejected inside staged CF (their effects - # would run once at trace time and freeze — a silent miscompile). - # - # This override lives ONLY on the PyIR subclass, which the DSL - # instantiates exclusively when pyir is enabled (see preprocess_mode). - # The base ``DSLPreprocessor.visit_With`` is untouched, so non-pyir - # compilation emits a byte-identical plain ``with`` — the wrapping is - # pyir-gated by construction, no runtime flag needed. Decorated / - # DSL-internal / stdlib managers pass through unchanged at run time. - visited = super().visit_With(node) - with_node = visited if isinstance(visited, ast.With) else None - if with_node is None: - return visited - for item in with_node.items: - ctx = item.context_expr - item.context_expr = ast.copy_location( - ast.Call( - func=_create_module_attribute( - "with_ctxmgr_check", - submodule_name="pyir_runtime", - lineno=getattr(ctx, "lineno", None), - col_offset=getattr(ctx, "col_offset", None), + lhs = ast.copy_location( + ast.IfExp( + test=ast.copy_location(test, node), + body=ast.Name(id=tmp_name, ctx=ast.Load()), + orelse=ast.Call( + func=_create_module_attribute( + helper_name, + use_base_dsl=False, + submodule_name=None, + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ + ast.Name(id=tmp_name, ctx=ast.Load()), + node.values[i], + ], + keywords=[], ), - args=[ctx], - keywords=[], ), - ctx, + node, ) - return with_node - def _visit_stmts_in_cf_scope_and_collect_definitions( - self, stmts: list[ast.stmt] - ) -> tuple[list[ast.stmt], set[str]]: - """Visit an isolated const_expr scope and return its definitions. + return ast.fix_missing_locations(lhs) - Used when visiting bodies of ``const_expr`` if/for/while — meta-level - control flow where only one branch executes at runtime. + def check_early_exit(self, tree: ast.AST, kind: str) -> None: + """ + Checks if a given region or scope in the provided Python code has early + exits. Local-catch awareness is a PyIR-mode fact: a ``raise`` in the + try suite of a region-local try/except is offered to that try's + handlers before it can unwind out, so it is not an early exit. + """ - The scope isolation prevents this bug:: + class EarlyExitChecker(ast.NodeVisitor): + def __init__(self, kind: str) -> None: + self.has_early_exit = False + self.early_exit_node: ast.AST | None = None + self.early_exit_type: str | None = None + self.kind = kind + self.loop_nest_level = 0 + self.guarded_try_depth = 0 + + def visit_Return(self, node: ast.Return) -> None: + self.has_early_exit = True + self.early_exit_node = node + self.early_exit_type = "return" + + def visit_Raise(self, node: ast.Raise) -> None: + # A raise in the try suite of a region-local try/except is + # offered to that try's handlers before it can unwind out. + if self.guarded_try_depth > 0: + return + self.has_early_exit = True + self.early_exit_node = node + self.early_exit_type = "raise" + + def visit_Try(self, node: ast.Try) -> None: + # Only the try suite is guarded; handler/else/finally bodies + # are not caught by this statement's own except clauses. + if node.handlers: + self.guarded_try_depth += 1 + for stmt in node.body: + self.visit(stmt) + if node.handlers: + self.guarded_try_depth -= 1 + for handler in node.handlers: + self.visit(handler) + for stmt in node.orelse: + self.visit(stmt) + for stmt in node.finalbody: + self.visit(stmt) + + def visit_Break(self, node: ast.Break) -> None: + if self.loop_nest_level == 0 and self.kind != "if": + self.has_early_exit = True + self.early_exit_node = node + self.early_exit_type = "break" + + def visit_Continue(self, node: ast.Continue) -> None: + if self.loop_nest_level == 0 and self.kind != "if": + self.has_early_exit = True + self.early_exit_node = node + self.early_exit_type = "continue" + + def visit_For(self, node: ast.For) -> None: + self.loop_nest_level += 1 + self.generic_visit(node) + self.loop_nest_level -= 1 + + def visit_While(self, node: ast.While) -> None: + self.loop_nest_level += 1 + self.generic_visit(node) + self.loop_nest_level -= 1 + + def visit_FunctionDef(self, node: ast.FunctionDef) -> None: + return + + checker = EarlyExitChecker(kind) + checker.generic_visit(tree) + if not checker.has_early_exit: + return + where = f"`{self.session_data.function_name}`" + ( + f" in `{self.session_data.class_name}`" + if self.session_data.class_name + else "" + ) + offender = checker.early_exit_node or tree + raise DSLUserCodeError( + DiagId.UNSUP_EARLY_EXIT, + filename=self.session_data.file_name, + lineno=getattr(offender, "lineno", None), + col_offset=getattr(offender, "col_offset", None), + end_col_offset=getattr(offender, "end_col_offset", None), + kind=checker.early_exit_type, + where=where, + ) - if const_expr(False): - cfg = make_a() # (1) adds 'cfg' to scope during AST visit - else: - cfg = make_b() # (2) without isolation, sees 'cfg' in scope - # → wraps with pyir_assign(pyir_read(cfg)) - # → UnboundLocalError at runtime because - # the if-branch never ran + def _visit_stmts_in_cf_scope( + self, stmts: list[ast.stmt], collect_bindings: "set[str] | None" = None + ) -> list[ast.stmt]: + """Visit statements in an isolated scope for const_expr branches. + + Used when visiting bodies of ``const_expr`` if/for/while — meta-level + control flow where only one branch executes at runtime. - With isolation, each branch gets its own scope. 'cfg' from (1) is - discarded before visiting (2), so (2) is correctly treated as a - first-definition (no pyir instrumentation). + The scope isolation gates INSTRUMENTATION only: a sibling branch must + not see this branch's first-definitions (it would wrap its own + first-def as a reassignment and pre-read an unbound name). The + BINDING fact itself survives — the caller collects this branch's + first-defs via *collect_bindings* and re-adds the union to the + enclosing scope after every arm is visited, so a later region's + write_args ∩ active_symbols keeps the carried name (Python truth: + the selected arm binds it). Outer-scope variables still work:: @@ -235,19 +961,28 @@ def _visit_stmts_in_cf_scope_and_collect_definitions( """ with self.session_data.scope_manager.enter_control_flow_scope(): result: list[ast.stmt] = [] - for stmt in stmts: - visited = self.visit(stmt) - if isinstance(visited, list): - result.extend(visited) - elif visited is not None: - result.append(visited) - definitions = set(self.session_data.scope_manager.scopes[-1]) - return result, definitions - - def _visit_stmts_in_cf_scope(self, stmts: list[ast.stmt]) -> list[ast.stmt]: - """Visit statements in an isolated const_expr scope.""" - result, _ = self._visit_stmts_in_cf_scope_and_collect_definitions(stmts) - return result + # Statement-insertion region for THIS body: hoists emitted while + # visiting these statements (walrus lowering, read/effect anchors) + # must land inside the const_expr body, not in the enclosing + # statement list where they would execute unconditionally (a + # never-taken arm's side effect) or before the loop (an unbound + # induction variable). + with Region(self.session_data, new_value=result): + for stmt in stmts: + visited = self.visit(stmt) + if isinstance(visited, list): + result.extend(visited) + elif visited is not None: + result.append(visited) + if collect_bindings is not None: + collect_bindings |= self.session_data.scope_manager.scopes[-1] + return result + + def _readd_constexpr_arm_bindings(self, born: "set[str]") -> None: + """R5c: after every arm of a const_expr statement is visited, its + arm-born bindings rejoin the enclosing scope as declared facts.""" + for name in born: + self.session_data.scope_manager.add_to_scope(name) def _handle_constexpr_for(self, node: ast.For) -> ast.For | list[ast.stmt]: """Override to add PyIR scope isolation for const_expr loops, and @@ -256,9 +991,25 @@ def _handle_constexpr_for(self, node: ast.For) -> ast.For | list[ast.stmt]: mutation guards know the body is trace-time-unrolled (not loop-carried) even when the enclosing CF is dynamic. """ - # Visit loop body in its own scope so first-definitions inside - # the body don't leak into the outer scope (PyIR only). - node.body = self._visit_stmts_in_cf_scope(node.body) + # A constexpr induction var is a fresh Meta int per iteration; skip-reference + # it so a name colliding with an earlier slot doesn't resurrect it. + induction_names = self._constexpr_induction_names(node.target) + already_skipped = { + name + for name in induction_names + if self.session_data.scope_manager.is_skip_reference_taking(name) + } + for name in induction_names: + self.session_data.scope_manager.add_skip_reference_taking(name) + born: set[str] = set() + try: + # Visit loop body in its own scope so first-definitions inside + # the body don't leak into the outer scope (PyIR only). + node.body = self._visit_stmts_in_cf_scope(node.body, collect_bindings=born) + finally: + for name in induction_names - already_skipped: + self.session_data.scope_manager.remove_skip_reference_taking(name) + self._readd_constexpr_arm_bindings(born - induction_names) # Wrap the unrolled body: enter; try: body; finally: exit. The # per-iteration bracket keeps the constexpr scope open inside the @@ -266,6 +1017,19 @@ def _handle_constexpr_for(self, node: ast.For) -> ast.For | list[ast.stmt]: node.body = self._wrap_body_in_constexpr_scope(node, node.body) return node + @staticmethod + def _constexpr_induction_names(target: ast.expr) -> set[str]: + """Collect the bare ``Name`` ids bound by a ``for`` target (tuple targets + included) so every rebound induction name is shielded.""" + names: set[str] = set() + if isinstance(target, ast.Name): + names.add(target.id) + elif isinstance(target, (ast.Tuple, ast.List)): + for elt in target.elts: + if isinstance(elt, ast.Name): + names.add(elt.id) + return names + # Names of the ``base_dsl.ast_helpers`` callbacks the preprocessor # emits to bracket a constexpr-governed loop/branch body. _ENTER_CONSTEXPR_LOOP = "enter_constexpr_loop" @@ -292,10 +1056,7 @@ def _call(name: str) -> ast.Expr: ast.Expr( value=ast.Call( func=_create_module_attribute( - name, - submodule_name="multi_stage_manager", - lineno=lineno, - col_offset=col_offset, + name, lineno=lineno, col_offset=col_offset ), args=[], keywords=[], @@ -322,7 +1083,7 @@ def _call(name: str) -> ast.Expr: # Function-name prefixes the preprocessor synthesizes for loops / branches. # These helper functions must NOT open a user-function scope -- they inherit - # their enclosing user function's scope so a loop-carried local (threaded + # their enclosing user function's scope so a loop-carried local (carried # through generated loop-body functions) stays on ONE slot key. _GENERATED_FN_PREFIXES = ( "loop_body_", @@ -337,66 +1098,536 @@ def _call(name: str) -> ast.Expr: "elif_region_", ) + @staticmethod + def _scope_cellvar_names(node: ast.AST) -> "list[str]": + """Names bound at *node*'s scope level and referenced inside a nested + scope (def / lambda / comprehension / class body): the frame's closure + cellvars. Over-approximation is transparent (the registration lambda + itself closes over the name); global-declared names resolve as globals + in the lambda and self-filter at registration.""" + bound: "set[str]" = set() + nested_refs: "set[str]" = set() + args = getattr(node, "args", None) + if args is not None: + for a in (*args.posonlyargs, *args.args, *args.kwonlyargs): + bound.add(a.arg) + if args.vararg is not None: + bound.add(args.vararg.arg) + if args.kwarg is not None: + bound.add(args.kwarg.arg) + _nested_kinds = ( + ast.FunctionDef, + ast.AsyncFunctionDef, + ast.Lambda, + ast.ClassDef, + ast.ListComp, + ast.SetComp, + ast.DictComp, + ast.GeneratorExp, + ) + + def _walk(n: ast.AST, in_nested: bool) -> None: + for child in ast.iter_child_nodes(n): + if isinstance(child, ast.Name): + if in_nested: + nested_refs.add(child.id) + elif isinstance(child.ctx, (ast.Store, ast.Del)): + bound.add(child.id) + elif not in_nested and isinstance( + child, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef) + ): + bound.add(child.name) + elif not in_nested and isinstance(child, ast.Global): + bound.difference_update(child.names) + elif not in_nested and isinstance(child, (ast.Import, ast.ImportFrom)): + for alias in child.names: + bound.add((alias.asname or alias.name).split(".")[0]) + # `except ... as name` and match capture patterns are binding + # forms too; their targets are str fields, not Store Names. + elif not in_nested and isinstance( + child, (ast.ExceptHandler, ast.MatchAs, ast.MatchStar) + ): + if child.name is not None: + bound.add(child.name) + elif not in_nested and isinstance(child, ast.MatchMapping): + if child.rest is not None: + bound.add(child.rest) + _walk(child, in_nested or isinstance(child, _nested_kinds)) + + for stmt in getattr(node, "body", []): + if isinstance(stmt, ast.Name): + continue + _walk(stmt, isinstance(stmt, _nested_kinds)) + if isinstance(stmt, (ast.FunctionDef, ast.AsyncFunctionDef, ast.ClassDef)): + bound.add(stmt.name) + return sorted(bound & nested_refs) + + def _scope_cell_registration_stmt( + self, node: "ast.FunctionDef" + ) -> "ast.Expr | None": + """The ``pyir_register_scope_cells(lambda: (c1, ..))`` entry statement + for *node*, or ``None`` when the scope owns no cells. The lambda's + ``__closure__`` exposes the frame's cell objects without evaluating + any cellvar (creation never touches the bindings). The scope's own + ``nonlocal`` names ride the same probe: their freevar cells ARE the + enclosing scope's cells, so registration exposes the shared identity.""" + cell_names = self._scope_cellvar_names(node) + for n in _own_scope_nonlocal_names(node.body): + if n not in cell_names: + cell_names.append(n) + if not cell_names: + return None + probe = ast.Lambda( + args=ast.arguments( + posonlyargs=[], + args=[], + kwonlyargs=[], + kw_defaults=[], + defaults=[], + ), + body=ast.Tuple( + elts=[ast.Name(id=n, ctx=ast.Load()) for n in cell_names], + ctx=ast.Load(), + ), + ) + stmt = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "pyir_register_scope_cells", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[probe], + keywords=[], + ) + ) + return ast.copy_location(ast.fix_missing_locations(stmt), node) + def _wrap_body_in_fn_scope( self, node: ast.stmt, body: list[ast.stmt] ) -> list[ast.stmt]: - """Bracket a USER function body with ``pyir_function_scope()`` / - ``pyir_function_scope()`` so frame-local slot keys are qualified by the owning + """Bracket a USER function body with ``pyir_enter_fn()`` / + ``pyir_exit_fn()`` so frame-local slot keys are qualified by the owning user function (preventing same-named locals in two different functions from colliding on one ``pyir.ref``). ```` is ``id(node)`` -- unique - per ``FunctionDef`` and stable for the preprocessing pass.""" + per ``FunctionDef`` and stable for the preprocessing pass. Generated + loop / branch helper functions are not wrapped (see ``visit_FunctionDef``) + so they inherit this scope at runtime.""" if not body: return body lineno = node.lineno col_offset = node.col_offset - withStmt = ast.With( - items=[ - ast.withitem( - context_expr=ast.Call( - func=_create_module_attribute( - "pyir_function_scope", - submodule_name="pyir_runtime", - lineno=lineno, - col_offset=col_offset, + def _call(name: str, args: list[ast.expr]) -> ast.Expr: + return ast.copy_location( + ast.fix_missing_locations( + ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + name, lineno=lineno, col_offset=col_offset + ), + args=args, + keywords=[], + ) + ) + ), + node, + ) + + # F-CEPLACE: parameters are DECLARED outside-born at scope entry (the + # seeding is the fact producer; absence is never defaulted). + args = getattr(node, "args", None) + param_names: "list[str]" = [] + if args is not None: + for a in (*args.posonlyargs, *args.args, *args.kwonlyargs): + param_names.append(a.arg) + if args.vararg is not None: + param_names.append(args.vararg.arg) + if args.kwarg is not None: + param_names.append(args.kwarg.arg) + entry_calls: "list[ast.stmt]" = [ + _call("pyir_enter_fn", [ast.Constant(value=id(node))]) + ] + if param_names: + entry_calls.append( + _call( + "pyir_seed_param_bindings", + [ast.Constant(value=n) for n in param_names], + ) + ) + # DECLARE the function's own ``nonlocal`` names on the scope (syntactic + # fact): a bare-name choke then recognizes closure-cell bindings. + nonlocal_names = _own_scope_nonlocal_names(body) + if nonlocal_names: + entry_calls.append( + _call( + "pyir_register_nonlocal_names", + [ast.Constant(value=n) for n in nonlocal_names], + ) + ) + # A parameter binding IS this scope's first-def: re-home a binding + # routed to another scope's live row onto this scope's own row. + if args is not None: + for a in (*args.posonlyargs, *args.args, *args.kwonlyargs): + entry_calls.append( + ast.copy_location( + ast.fix_missing_locations( + ast.Assign( + targets=[ast.Name(id=a.arg, ctx=ast.Store())], + value=ast.Call( + func=_create_runtime_attribute( + "pyir_bind_param", + lineno=lineno, + col_offset=col_offset, + ), + args=[ + ast.Constant(value=a.arg), + ast.Name(id=a.arg, ctx=ast.Load()), + ], + keywords=[], + ), + ) ), - args=[ast.Constant(value=id(node))], - keywords=[], - ), - optional_vars=None, + node, + ) ) - ], - body=body, - ) - return [ast.copy_location(ast.fix_missing_locations(withStmt), node)] + return entry_calls + [ + ast.copy_location( + ast.fix_missing_locations( + ast.Try( + body=body, + handlers=[], + orelse=[], + finalbody=[_call("pyir_exit_fn", [])], + ) + ), + node, + ), + ] + + # >0 while visiting a def nested inside the outermost one: a nested def + # compiles from this same instrumented AST, so the dispatcher fast-exits it. + _pyir_fn_def_depth: int = 0 @override - def visit_FunctionDef(self, node: ast.FunctionDef) -> ast.FunctionDef: - """Wrap USER function bodies so local slot keys are unique per user function. - Generated loop / branch helper functions are skipped -- they inherit the - enclosing user scope.""" - result = super().visit_FunctionDef(node) + def visit_FunctionDef( + self, node: ast.FunctionDef + ) -> "ast.FunctionDef | list[ast.stmt]": + """Wrap USER function bodies with ``pyir_enter_fn`` / ``pyir_exit_fn`` so + local slot keys are unique per user function. Generated loop / branch + helper functions are skipped -- they inherit the enclosing user scope.""" + _saved_class_depth = self._pyir_class_body_depth + self._pyir_class_body_depth = 0 + self._pyir_fn_def_depth += 1 + self._pyir_import_scope_stack = self._pyir_import_scope_stack + ( + self._scope_import_bindings(node), + ) + # This def's whole body is one scope of the nonlocal owner table. + self._function_scope_stack = self._function_scope_stack + (node,) + try: + result = super().visit_FunctionDef(node) + finally: + self._pyir_fn_def_depth -= 1 + self._pyir_class_body_depth = _saved_class_depth + self._pyir_import_scope_stack = self._pyir_import_scope_stack[:-1] + self._function_scope_stack = self._function_scope_stack[:-1] + if isinstance(result, ast.FunctionDef): + # Register the frame's closure cells at entry (user AND generated + # fns own cells); generated region fns register on the inherited + # scope, matching where their locals' slot keys resolve. + _cell_reg = self._scope_cell_registration_stmt(result) + if _cell_reg is not None: + result.body.insert(0, _cell_reg) + # Region-synthesized defs (created by the CF outliners, never + # revisited here) own the carry-param cells nested closures + # capture; register them so cell identity resolves to the row. + for _fd in ast.walk(result): + if ( + isinstance(_fd, ast.FunctionDef) + and getattr(_fd, "_pyir_synth_scope", False) + and not getattr(_fd, "_pyir_cells_registered", False) + ): + _fd._pyir_cells_registered = True # type: ignore[attr-defined] + _synth_reg = self._scope_cell_registration_stmt(_fd) + if _synth_reg is not None: + _fd.body.insert(0, _synth_reg) if isinstance(result, ast.FunctionDef) and not result.name.startswith( self._GENERATED_FN_PREFIXES ): result.body = self._wrap_body_in_fn_scope(result, result.body) + # A nested user def's writes already commit through the instrumented + # chokes; mark its runtime object so the dispatcher fast-exits it. + # The stamp applies as the INNERMOST decorator so it attests exactly + # the raw function object compiled from this instrumented AST — a + # decorator wrapper around it is never attested (its uninstrumented + # writes stay visible to the boundary observer), and an + # undecoratable decoration result (e.g. ``@property``) never + # receives a post-decoration attribute store. + if self._pyir_fn_def_depth > 0: + stamp = _create_runtime_attribute( + "_pyir_stamp_rewritten", + lineno=result.lineno, + col_offset=result.col_offset, + ) + result.decorator_list.append(stamp) + ast.fix_missing_locations(result) return result - def _prepare_loop_induction_var(self, node: ast.For) -> None: - """Override to mark induction variable for skip instrumentation. + @override + def _post_visit_function_body(self, func_def: ast.FunctionDef) -> "list[ast.stmt]": + """Bind the runtime-module alias once at function entry (every + emitted choke call resolves through it; nested regions close over + it), then run the call-boundary wrap pass; when any call was + wrapped, also bind the dispatcher alias the same way.""" + stmts: "list[ast.stmt]" = [ + ast.Assign( + targets=[ast.Name(id=_PYIR_RUNTIME_ALIAS, ctx=ast.Store())], + value=_create_module_attribute("pyir_runtime", submodule_name=None), + ) + ] + if _CallBoundaryWrapPass().run_on_function_body(func_def): + stmts.append( + ast.Assign( + targets=[ + ast.Name(id=_CallBoundaryWrapPass.ALIAS_NAME, ctx=ast.Store()) + ], + value=_create_runtime_attribute("_pyir_call_boundary_"), + ) + ) + return stmts + + # >0 while visiting a CLASS body's immediate statements: an ``AnnAssign`` + # there is a field declaration and must not decompose to an assignment. + _pyir_class_body_depth: int = 0 - Mark the induction variable so _handle_negative_step's injected - `idx = offset - idx if isNeg else idx` is not instrumented. - Creating a reference for the block argument would cause issues. - """ + # One (bound, plain-import-only, from-import-only) name frame per function + # scope being visited; import-only names classify as module spellings, + # from-import-only names as symbol-rooted reads (the global-read contract). + _pyir_import_scope_stack: "tuple[tuple[frozenset, frozenset, frozenset], ...]" = () + + @staticmethod + def _scope_import_bindings( + node: ast.AST, + ) -> "tuple[frozenset, frozenset, frozenset]": + """(names bound at *node*'s scope level, the subset bound ONLY by plain + ``import`` statements, the subset bound ONLY by ``from`` imports). + Every plain-import binding is a module object (LangRef 3.12 section + 7.11), so those names classify as module spellings for the read chokes; + a from-import binding is a symbol read off its source module, so its + attr reads classify like function-scope global-object reads. Any other + binding construct on the name disqualifies it (over-approximation of + "other" is transparent: the name keeps its ordinary-local class).""" + imported: "set[str]" = set() + from_imported: "set[str]" = set() + other: "set[str]" = set() + args = getattr(node, "args", None) + if args is not None: + for a in (*args.posonlyargs, *args.args, *args.kwonlyargs): + other.add(a.arg) + if args.vararg is not None: + other.add(args.vararg.arg) + if args.kwarg is not None: + other.add(args.kwarg.arg) + _nested_scopes = ( + ast.FunctionDef, + ast.AsyncFunctionDef, + ast.Lambda, + ast.ClassDef, + ) + + def _walk(n: ast.AST) -> None: + for child in ast.iter_child_nodes(n): + if isinstance(child, ast.Import): + for alias in child.names: + imported.add(alias.asname or alias.name.split(".")[0]) + elif isinstance(child, ast.ImportFrom): + for alias in child.names: + from_imported.add(alias.asname or alias.name) + elif isinstance(child, ast.Name) and isinstance( + child.ctx, (ast.Store, ast.Del) + ): + other.add(child.id) + elif isinstance(child, (ast.Global, ast.Nonlocal)): + other.update(child.names) + elif isinstance(child, ast.ExceptHandler) and child.name: + other.add(child.name) + elif isinstance(child, (ast.MatchAs, ast.MatchStar)) and child.name: + other.add(child.name) + elif isinstance(child, ast.MatchMapping) and child.rest: + other.add(child.rest) + elif isinstance(child, _nested_scopes): + name = getattr(child, "name", None) + if name is not None: + other.add(name) + continue # a nested scope's own bindings stay its own + _walk(child) + + _walk(node) + return ( + frozenset(imported | from_imported | other), + frozenset(imported - from_imported - other), + frozenset(from_imported - imported - other), + ) + + def _pyir_from_import_bound(self, name: str) -> bool: + """True when *name*, in its nearest binding scope, is bound ONLY by + ``from``-import statements: its attr reads take the symbol-rooted + record-only choke (the function-scope global-object read contract).""" + for bound, _imported, from_imported in reversed(self._pyir_import_scope_stack): + if name in from_imported: + return True + if name in bound: + return False + return False + + def visit_Import(self, node: ast.Import) -> "ast.stmt | list[ast.stmt]": + """Bracket an in-body ``import`` with the staged-CF wall and the F-SPEC + record arm: the import executes once at trace (LangRef 3.12 section + 7.11), so inside dynamic staged CF the wall refuses BEFORE the module's + side effects run; in meta flow every bound name records a verified + spec root (a plain import binds only modules).""" + pairs = [ + (alias.name, alias.asname or alias.name.split(".")[0]) + for alias in node.names + ] + return self._pyir_wrap_import(node, None, 0, pairs) + + def visit_ImportFrom(self, node: ast.ImportFrom) -> "ast.stmt | list[ast.stmt]": + """``from``-import twin of :py:meth:`visit_Import`: the record arm + re-derives each bound value off the resolved source module (scalars + pin value rows, objects root for their leaf reads).""" + pairs = [(alias.name, alias.asname or alias.name) for alias in node.names] + return self._pyir_wrap_import(node, node.module, node.level, pairs) + + def _pyir_reading_package(self) -> "str | None": + """The defining module's declared package (``__package__``, falling + back to ``__spec__.parent`` -- LangRef 3.12 section 5.4.2): a + rewrite-time constant anchoring relative-import resolution.""" + g = self.session_data.function_globals + if not g: + return None + pkg = g.get("__package__") + if isinstance(pkg, str): + return pkg + parent = getattr(g.get("__spec__"), "parent", None) + return parent if isinstance(parent, str) else None + + def _pyir_wrap_import( + self, + node: ast.stmt, + module: "str | None", + level: int, + pairs: "list[tuple[str, str]]", + ) -> "ast.stmt | list[ast.stmt]": + """[guard, import, record] bracket for one import statement; the native + statement itself is untouched (Python truth for binding semantics).""" + if getattr(node, "_pyir_import_wrapped", False): + return node + node._pyir_import_wrapped = True # type: ignore[attr-defined] + package = self._pyir_reading_package() if level else None + guard: ast.stmt = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_import_guard", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ + ast.Constant(value=ast.unparse(node)), + ast.Constant(value=module), + ast.Constant(value=level), + ast.Constant(value=package), + *(ast.Constant(value=public) for public, _bound in pairs), + ], + keywords=[], + ) + ) + record: ast.stmt = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_import_record", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ + ast.Constant(value=module), + ast.Constant(value=level), + ast.Constant(value=package), + *( + ast.Tuple( + elts=[ + ast.Constant(value=public), + ast.Name(id=bound, ctx=ast.Load()), + ], + ctx=ast.Load(), + ) + for public, bound in pairs + ), + ], + keywords=[], + ) + ) + return [ + ast.copy_location(ast.fix_missing_locations(guard), node), + node, + ast.copy_location(ast.fix_missing_locations(record), node), + ] + + def visit_ClassDef(self, node: ast.ClassDef) -> ast.AST: + """Run a class body in its own PyIR scope so a field annotation does not + leak into the enclosing function scope and mask a local's first-def.""" + self._pyir_class_body_depth += 1 + try: + with self.session_data.scope_manager.enter_local_scope(): + return super().visit_ClassDef(node) + finally: + self._pyir_class_body_depth -= 1 + + def visit_With(self, node: ast.With) -> ast.AST: + """Route each context expression through the ctx boundary: ``with`` + runs its dunders outside ``ast.Call``, invisible to the call wrap.""" + result = super().visit_With(node) # scope registration + child visits + if isinstance(result, ast.With): + for item in result.items: + item.context_expr = ast.copy_location( + ast.Call( + func=_create_runtime_attribute( + "_pyir_ctx_boundary", + lineno=getattr(item.context_expr, "lineno", 0), + col_offset=getattr(item.context_expr, "col_offset", 0), + ), + args=[item.context_expr], + keywords=[], + ), + item.context_expr, + ) + ast.fix_missing_locations(result) + return result + + def _prepare_loop_induction_var( + self, + node: ast.For, + target_is_live_after_loop: bool = False, + loop_carried_var_name: str | None = None, + ) -> None: + """Skip instrumentation of the induction target (a ref on the block + argument would not dominate), and record the synthetic live-out carry.""" if isinstance(node.target, ast.Name): self.session_data.scope_manager.add_skip_reference_taking(node.target.id) + if target_is_live_after_loop and loop_carried_var_name is not None: + # Write-only in the body, so record it for the prologue's forced + # promotion and scope it so the body write stores into that cell. + self.session_data.liveout_loop_carried_vars[id(node)] = ( + loop_carried_var_name + ) + self.session_data.scope_manager.add_to_scope(loop_carried_var_name) def _cleanup_loop_induction_var(self, node: ast.For) -> None: """Override to remove skip flag after loop function creation.""" if isinstance(node.target, ast.Name): self.session_data.scope_manager.remove_skip_reference_taking(node.target.id) + self.session_data.liveout_loop_carried_vars.pop(id(node), None) def _is_element_in_scope(self, elt: ast.expr) -> bool: """Check if a single tuple element needs pyir instrumentation.""" @@ -431,7 +1662,7 @@ def _is_target_in_scope(self, target: ast.expr) -> bool: - First-time definitions (not yet in scope) - Variables in skip_pyir_reference_taking (e.g. loop induction variables) """ - if isinstance(target, ast.Tuple): + if isinstance(target, (ast.Tuple, ast.List)): return any(self._is_element_in_scope(elt) for elt in target.elts) if isinstance(target, ast.Name): if self.session_data.scope_manager.is_skip_reference_taking(target.id): @@ -498,7 +1729,7 @@ def _is_subscript_dict_style(self, node: ast.Subscript) -> bool: Only string-constant keys are supported. Integer/variable subscripts (tensor[i], array[idx]) are genuine memory ops handled by MLIR and do - not need PyIR SSA threading. + not need PyIR SSA carry. """ if ( isinstance(node.slice, ast.Constant) @@ -510,69 +1741,523 @@ def _is_subscript_dict_style(self, node: ast.Subscript) -> bool: return False @staticmethod - def _names_read_before_and_after_walrus(node: ast.AST) -> set[str]: - """Names ``X`` read BOTH before AND after the same-statement walrus - ``(X := ...)`` (outside the walrus subtree). Such a name is - UN-ANCHORABLE: element0 wants the pre-walrus value and a later element - wants the post-walrus value from ONE ``X``. ``a, b, c = m1, (m1 := - m1+1), m1`` silently miscompiled the mismatched read (verified on this - base), so these shapes are refused loudly instead. Self-contained AST - scan (same-statement, positional, walrus-RHS reads excluded).""" - walrus_pos: dict[str, tuple[int, int]] = {} - for sub in ast.walk(node): - if isinstance(sub, ast.NamedExpr) and isinstance(sub.target, ast.Name): - pos = (sub.lineno, sub.col_offset) - tgt = sub.target.id - if tgt not in walrus_pos or pos < walrus_pos[tgt]: - walrus_pos[tgt] = pos - if not walrus_pos: - return set() - loads: dict[str, list[tuple[int, int]]] = {} - - def _scan(n: ast.AST, in_walrus: bool) -> None: - child_in_walrus = in_walrus or isinstance(n, ast.NamedExpr) - if ( - isinstance(n, ast.Name) - and isinstance(n.ctx, ast.Load) - and n.id in walrus_pos - and not in_walrus + def _walrus_stored_names(value: ast.expr) -> list[str]: + """Names walrus'd anywhere in a statement's value expression + (lambda subtrees are opaque: their walrus binds lambda-locally).""" + names: list[str] = [] + _scan_stack: list[ast.AST] = [value] + while _scan_stack: + n = _scan_stack.pop() + if isinstance(n, ast.Lambda): + continue + if isinstance(n, ast.NamedExpr) and isinstance(n.target, ast.Name): + if n.target.id not in names: + names.append(n.target.id) + _scan_stack.extend(ast.iter_child_nodes(n)) + return names + + def _hoist_pre_walrus_reads(self, node: ast.stmt, value: "ast.expr | None") -> None: + """Statement-top read anchors for pre-walrus reads (Python evaluation order).""" + if ( + value is None + or self._pyir_lambda_depth + or not self.session_data.region_stack + ): + return + walrus_names = self._walrus_stored_names(value) + if not walrus_names: + return + anchor_by_read_id: dict[int, str] = {} + anchor_stmts: list[ast.stmt] = [] + for name in walrus_names: + reads: list[ast.Name] = [] + + def _walk(n: ast.AST) -> bool: + """DFS in field order (Python evaluation order for the + expression forms allowed here); True = walrus reached.""" + if isinstance(n, ast.Lambda): + return False + if ( + isinstance(n, ast.NamedExpr) + and isinstance(n.target, ast.Name) + and n.target.id == name + ): + return True + if ( + isinstance(n, ast.Name) + and isinstance(n.ctx, ast.Load) + and n.id == name + ): + reads.append(n) + return False + for child in ast.iter_child_nodes(n): + if _walk(child): + return True + return False + + _walk(value) + if not reads or not self._is_target_in_scope( + ast.Name(id=name, ctx=ast.Store()) ): - loads.setdefault(n.id, []).append((n.lineno, n.col_offset)) + continue + anchor = f"_pyir_wpre_{self.session_data.counter}" + self.session_data.counter += 1 + anchor_call = ast.Call( + func=_create_runtime_attribute( + "_pyir_anchor_statement_read", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ast.Name(id=name, ctx=ast.Load())], + keywords=[], + ) + anchor_assign = ast.Assign( + targets=[ast.Name(id=anchor, ctx=ast.Store())], + value=anchor_call, + ) + anchor_stmts.append( + ast.copy_location(ast.fix_missing_locations(anchor_assign), node) + ) + for r in reads: + anchor_by_read_id[id(r)] = anchor + if not anchor_stmts: + return + + class _RewritePreWalrusReads(ast.NodeTransformer): + def visit_Name(self, n: ast.Name) -> ast.Name: + repl = anchor_by_read_id.get(id(n)) + if repl is not None: + return ast.copy_location(ast.Name(id=repl, ctx=ast.Load()), n) + return n + + def visit_Lambda(self, n: ast.Lambda) -> ast.Lambda: + return n # opaque + + _RewritePreWalrusReads().visit(node) + ast.fix_missing_locations(node) + # Before the walrus's own hoisted assign (visit_NamedExpr appends + # during the statement's generic_visit, which runs after this). + self.session_data.region_stack[-1].append_new_stmts(anchor_stmts) + + def _hoist_pre_walrus_augassign_base(self, node: ast.AugAssign) -> "str | None": + """Statement-top anchor for the implicit base read of ``name op= value`` + when *value* walrus-rebinds ``name``. Python loads the base BEFORE + evaluating the RHS, but the walrus lowers to an assign hoisted ahead of + the statement, so a statement-position base read would see the + post-walrus value. Returns the anchor temp for + ``_insert_pyir_augassign`` to consume as the in-place-op base; the + statement-position ``pyir_read`` refresh still runs (staged-slot load + semantics for the store's own old-value capture).""" + if ( + self._pyir_lambda_depth + or not self.session_data.region_stack + or not isinstance(node.target, ast.Name) + or not self._is_target_in_scope(node.target) + or node.target.id not in self._walrus_stored_names(node.value) + ): + return None + anchor = f"_pyir_wbase_{self.session_data.counter}" + self.session_data.counter += 1 + anchor_call = ast.Call( + func=_create_runtime_attribute( + "_pyir_anchor_statement_read", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ast.Name(id=node.target.id, ctx=ast.Load())], + keywords=[], + ) + anchor_assign = ast.Assign( + targets=[ast.Name(id=anchor, ctx=ast.Store())], + value=anchor_call, + ) + # Before the walrus's own hoisted assign (visit_NamedExpr appends + # during the statement's generic_visit, which runs after this). + self.session_data.region_stack[-1].append_new_stmts( + [ast.copy_location(ast.fix_missing_locations(anchor_assign), node)] + ) + return anchor + + def _hoist_pre_walrus_effects( + self, node: ast.stmt, value: "ast.expr | None" + ) -> None: + """Statement-top effect anchors (Python evaluation order): a call that + completes BEFORE the first walrus in *value* is hoisted into a temp, in + completion order, so the walrus lowering's own hoist cannot move it + after the walrus's write. A call the anchor cannot represent + order-faithfully — one in a conditionally-evaluated position before a + walrus, or one completing between two walruses — refuses loudly + (never silently reordered).""" + if ( + value is None + or self._pyir_lambda_depth + or not self.session_data.region_stack + ): + return + # Effect units in completion (post-DFS) order == Python effect order. + events: "list[tuple[str, ast.AST, bool]]" = [] + + def _contains_call(n: ast.AST) -> bool: + return any(isinstance(sub, ast.Call) for sub in ast.walk(n)) + + def _scan(n: ast.AST, conditional: bool) -> None: + if isinstance(n, (ast.Lambda, ast.GeneratorExp)): + return # deferred execution: not this statement's effects + if isinstance(n, (ast.ListComp, ast.SetComp, ast.DictComp)): + # Eager comprehension: one opaque effect unit when it can call. + if _contains_call(n): + events.append(("call", n, conditional)) + return + if isinstance(n, ast.BoolOp): + _scan(n.values[0], conditional) + for v in n.values[1:]: + _scan(v, True) + return + if isinstance(n, ast.IfExp): + _scan(n.test, conditional) + _scan(n.body, True) + _scan(n.orelse, True) + return for child in ast.iter_child_nodes(n): - _scan(child, child_in_walrus) + _scan(child, conditional) + if isinstance(n, ast.NamedExpr): + events.append(("walrus", n, conditional)) + elif isinstance(n, ast.Call): + events.append(("call", n, conditional)) + + _scan(value, False) + walrus_positions = [i for i, (k, _, _) in enumerate(events) if k == "walrus"] + if not walrus_positions: + return + first_w, last_w = walrus_positions[0], walrus_positions[-1] + + def _refuse(call_node: ast.AST, walrus_node: ast.AST) -> None: + target = getattr(walrus_node, "target", None) + raise DSLUserCodeError( + DiagId.UNSUP_WALRUS_EVAL_ORDER, + filename=self.session_data.file_name, + lineno=getattr(call_node, "lineno", None), + col_offset=getattr(call_node, "col_offset", None), + end_col_offset=getattr(call_node, "end_col_offset", None), + call=_unparse_safe(call_node), + name=target.id + if isinstance(target, ast.Name) + else _unparse_safe(walrus_node), + ) - _scan(node, False) - both: set[str] = set() - for name, w_pos in walrus_pos.items(): - positions = loads.get(name) - if ( - positions - and any(p < w_pos for p in positions) - and any(p > w_pos for p in positions) + anchors: "list[ast.AST]" = [] + for i, (kind, n, conditional) in enumerate(events): + if kind != "call": + continue + if i < first_w: + if conditional: + _refuse(n, events[first_w][1]) + anchors.append(n) + elif i < last_w: + # Completes between two walruses: the lowering would run it + # after every hoisted walrus assign. + nxt = next(j for j in walrus_positions if j > i) + _refuse(n, events[nxt][1]) + if not anchors: + return + # Only outermost units hoist; an inner call evaluates inside its + # enclosing anchored unit, so order is preserved within. + outer = [ + c + for c in anchors + if not any( + c2 is not c and any(sub is c for sub in ast.walk(c2)) for c2 in anchors + ) + ] + replace_by_id: "dict[int, str]" = {} + for c in outer: + temp = f"_pyir_weff_{self.session_data.counter}" + self.session_data.counter += 1 + assign = ast.copy_location( + ast.Assign( + targets=[ast.copy_location(ast.Name(id=temp, ctx=ast.Store()), c)], + value=c, # type: ignore[arg-type] + ), + node, + ) + ast.fix_missing_locations(assign) + # The synthetic temp is machinery-owned (no pyir_assign choke), + # but the hoisted VALUE still visits: its reads and call-boundary + # instrumentation execute at the anchor position, as Python does. + self.session_data.scope_manager.add_skip_reference_taking(temp) + visited = self.visit(assign) + self.session_data.region_stack[-1].append_new_stmts( + visited if isinstance(visited, list) else [visited] + ) + replace_by_id[id(c)] = temp + + class _RewriteAnchoredEffects(ast.NodeTransformer): + def visit(self, n: ast.AST) -> ast.AST: + repl = replace_by_id.get(id(n)) + if repl is not None: + return ast.copy_location(ast.Name(id=repl, ctx=ast.Load()), n) + return super().visit(n) + + _RewriteAnchoredEffects().visit(node) + ast.fix_missing_locations(node) + + # ------------------------------------------------------------------ + # Walrus (``ast.NamedExpr``) lowering + # ------------------------------------------------------------------ + # >0 while visiting an ``ast.Lambda`` body: a walrus there binds in the LAMBDA's own + # scope at call time (PEP 572), so it is not this statement's write and must be left + _pyir_lambda_depth: int = 0 + + @staticmethod + def _find_eager_walrus(root: ast.AST) -> "ast.NamedExpr | None": + """First ``ast.NamedExpr`` under *root* whose write would land in the ENCLOSING + function scope when *root* evaluates.""" + stack: list[ast.AST] = [root] + while stack: + n = stack.pop() + if isinstance(n, ast.Lambda): + continue + if isinstance(n, ast.NamedExpr): + return n + stack.extend(ast.iter_child_nodes(n)) + return None + + def _refuse_conditional_walrus(self, sub: "ast.expr | None") -> None: + """Curated refusal for a walrus in a conditionally-evaluated position.""" + if sub is None: + return + walrus = self._find_eager_walrus(sub) + if walrus is None: + return + raise DSLUserCodeError( + DiagId.UNSUP_WALRUS_CONDITIONAL, + filename=self.session_data.file_name, + lineno=getattr(walrus, "lineno", None), + col_offset=getattr(walrus, "col_offset", None), + end_col_offset=getattr(walrus, "end_col_offset", None), + name=walrus.target.id + if isinstance(walrus.target, ast.Name) + else ast.unparse(walrus.target), + ) + + def visit_BoolOp(self, node: ast.BoolOp) -> ast.expr: + # ``values[1:]`` of and/or are conditionally evaluated; ``values[0]`` + # always evaluates, so a walrus there lowers normally. + for v in node.values[1:]: + self._refuse_conditional_walrus(v) + return self._expand_boolop_evaluate_once(node) + + def visit_IfExp(self, node: ast.IfExp) -> ast.Call: + # Ternary arms evaluate conditionally; the test always evaluates. + self._refuse_conditional_walrus(node.body) + self._refuse_conditional_walrus(node.orelse) + return super().visit_IfExp(node) + + def visit_Lambda(self, node: ast.Lambda) -> ast.Lambda: + self._pyir_lambda_depth += 1 + try: + return super().visit_Lambda(node) + finally: + self._pyir_lambda_depth -= 1 + + def _visit_Comprehension( + self, node: "_ComprehensionT", ele_visitor: "Callable[..., Any]" + ) -> "_ComprehensionT": + # A comprehension walrus writes the ENCLOSING function scope from the + # comprehension's looped scope (PEP 572): iteration-dependent, refuse. + self._refuse_conditional_walrus(node) + return super()._visit_Comprehension(node, ele_visitor) + + def visit_NamedExpr(self, node: ast.NamedExpr) -> ast.expr: + """Lower ``name := value`` into a real, fully instrumented assignment.""" + if self._pyir_lambda_depth or not self.session_data.region_stack: + self.generic_visit(node) + return node + # PEP 572: a NamedExpr target is always a plain name. + assert isinstance(node.target, ast.Name) + name = node.target.id + assign = ast.copy_location( + ast.Assign( + targets=[ast.copy_location(ast.Name(id=name, ctx=ast.Store()), node)], + value=node.value, + ), + node, + ) + ast.fix_missing_locations(assign) + visited = self.visit(assign) + self.session_data.region_stack[-1].append_new_stmts( + visited if isinstance(visited, list) else [visited] + ) + return ast.copy_location(ast.Name(id=name, ctx=ast.Load()), node) + + def _nested_unpack_preread(self, name: str, node: ast.Assign) -> ast.stmt: + """Guarded ``name = pyir_read(name, name)`` (the flat Step-1 shape): + re-binds the name to its ref-loaded value BEFORE the original RHS + evaluates, since the follow-up's own Step-1 read fires only after.""" + read_call = ast.Call( + func=_create_runtime_attribute( + "pyir_read", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ + ast.Constant(value=name), + ast.Name(id=name, ctx=ast.Load()), + ], + keywords=[], + ) + tried = ast.Try( + body=[ + ast.Assign( + targets=[ast.Name(id=name, ctx=ast.Store())], + value=read_call, + ) + ], + handlers=[ + ast.ExceptHandler( + type=ast.Tuple( + elts=[ + ast.Name(id="NameError", ctx=ast.Load()), + ast.Name(id="UnboundLocalError", ctx=ast.Load()), + ], + ctx=ast.Load(), + ), + name=None, + body=[ast.Pass()], + ) + ], + orelse=[], + finalbody=[], + ) + return ast.copy_location(ast.fix_missing_locations(tried), node) + + @staticmethod + def _is_multi_targets_assign(node: ast.Assign) -> bool: + """``t1 = t2 = rhs``: more than one target shares one RHS.""" + return len(node.targets) > 1 + + @staticmethod + def _is_nested_unpack_target(node: ast.Assign) -> bool: + """``a, (b, c) = rhs``: some element of a sequence target is itself a + sequence, which is what opens the extra unpack level.""" + return any( + isinstance( + elt.value if isinstance(elt, ast.Starred) else elt, + (ast.Tuple, ast.List), + ) + for target in node.targets + if isinstance(target, (ast.Tuple, ast.List)) + for elt in target.elts + ) + + def _split_nested_unpack_target(self, node: ast.Assign) -> "list[ast.stmt]": + """Split nested unpack targets so every level goes through the flat + tuple decompose: ``a, (b, c) = rhs`` becomes + ``_pyir_nest_0, _pyir_nest_1 = rhs`` followed by ``a = _pyir_nest_0`` + and ``b, c = _pyir_nest_1``. A nested Tuple/List element otherwise + unpacks bare -- no pyir_read/pyir_assign -- so its stores never reach + the ledger and staged loops carry frozen values. Python evaluates the + whole RHS before any element store, so the split preserves swap + pinning; deeper nesting recurses through the re-visit. + + EVERY element of EVERY sequence target is staged, not just the nested + ones, so the head statement makes no user-visible store and the + follow-ups replay the whole target list in source order. Python stores + targets strictly left to right while a follow-up runs after the head + statement, so an element left behind would store ahead of the staged + ones -- inverting the last writer when the two name the same place: + ``(x, y), x = (1, 2), 3`` must leave ``x == 3``, not ``1``. Staging + only the colliding elements is not an option: ``p = o`` makes + ``(o.a, y), p.a = ...`` name one place through two syntactically + distinct elements, so the collision is not decidable here. + + Returns the rewritten statements, already visited. The prereads + re-bind every in-scope Name before the RHS (flat + Step-1 ordering): a staged element no longer takes the head statement's + Step-1 read, so it would otherwise lose that read entirely. Follow-up + statements are marked so their recursive split skips redundant prereads + (their RHS is the already-pinned temp).""" + is_followup = getattr(node, "_pyir_nest_followup", False) + prereads: list[ast.stmt] = [] + preread_names: set[str] = set() + + def _unstar(elt: ast.expr) -> ast.expr: + return elt.value if isinstance(elt, ast.Starred) else elt + + def _collect_prereads(elt: ast.expr) -> None: + inner = _unstar(elt) + if isinstance(inner, (ast.Tuple, ast.List)): + for e in inner.elts: + _collect_prereads(e) + elif ( + isinstance(inner, ast.Name) + and inner.id != "_" + and inner.id not in preread_names + and self._is_element_in_scope(inner) ): - both.add(name) - return both + preread_names.add(inner.id) + prereads.append(self._nested_unpack_preread(inner.id, node)) - def visit_Assign(self, node: ast.Assign) -> ast.stmt | list[ast.stmt]: - """Override to add PyIR instrumentation for reassignments.""" - # A walrus-target name read on BOTH sides of its walrus in a tuple RHS - # is un-anchorable and silently miscompiled the mismatched read (fresh- - # Name AND subscript targets). Scoped to a Tuple target so the - # positional before/after test matches left-to-right RHS evaluation (a - # ternary/boolop RHS has positional != eval order). - _walrus_both = self._names_read_before_and_after_walrus(node) - if _walrus_both: - for _tgt in node.targets: - if isinstance(_tgt, ast.Tuple): - raise DSLUserCodeError( - DiagId.UNSUP_WALRUS_TUPLE_REBIND, - filename=self.session_data.file_name, - lineno=node.lineno, - col_offset=node.col_offset, - end_col_offset=getattr(node, "end_col_offset", None), - var=sorted(_walrus_both)[0], - ) - # Check scope BEFORE _visit_target adds the target. + followups: list[ast.Assign] = [] + for target in node.targets: + if not isinstance(target, (ast.Tuple, ast.List)): + continue + for i in range(len(target.elts)): + elt = target.elts[i] + inner = _unstar(elt) + if not is_followup: + _collect_prereads(elt) + tmp = f"_pyir_nest_{self.session_data.counter}" + self.session_data.counter += 1 + # The synthetic temp is machinery-owned: no first-def + # pyir_assign choke, no carry. + self.session_data.scope_manager.add_skip_reference_taking(tmp) + tmp_store: ast.expr = ast.Name(id=tmp, ctx=ast.Store()) + if isinstance(elt, ast.Starred): + tmp_store = ast.Starred(value=tmp_store, ctx=ast.Store()) + target.elts[i] = tmp_store + # A List element unpacks identically to a Tuple; normalize so + # the follow-up hits the tuple decompose. A non-nested element + # (star-stripped Name / Attribute / Subscript) is already its + # own follow-up target. + follow_target = ( + ast.Tuple(elts=inner.elts, ctx=ast.Store()) + if isinstance(inner, ast.List) + else inner + ) + followup = ast.Assign( + targets=[follow_target], + value=ast.Name(id=tmp, ctx=ast.Load()), + ) + followup._pyir_nest_followup = True # type: ignore[attr-defined] + followups.append(followup) + # ``_is_nested_unpack_target`` gated the call, so at least one element + # was staged and ``followups`` is never empty here. + result: list[ast.stmt] = list(prereads) + for stmt in [node, *followups]: + visited = self.visit( + ast.copy_location(ast.fix_missing_locations(stmt), node) + ) + result.extend(visited if isinstance(visited, list) else [visited]) + return result + + def _collect_assign_facts(self, node: ast.Assign) -> "list[ast.expr]": + """Collect the STATEMENT-level facts the exits below need, BEFORE any + target is visited. + + Returns ``targets_to_instrument`` (the dispatch needs it) and ATTACHES + the other two facts to the AST, where the exit that consumes each can + reach it from *node* alone: the first-def target list on the statement, + and each sequence target's in-scope element indices on that target. + Both have to be captured here and cannot be recomputed later -- + ``_visit_target`` adds first-definitions to scope right after this + call, which would make every element look in-scope. + + Order matters twice over: every fact must be read before + ``_visit_target`` runs, and ``_is_target_in_scope`` records + ``__init__`` attributes in ``seen_init_attrs`` as a side effect -- so + the blocks below keep their relative order. + """ # This distinguishes first-time definitions from reassignments. targets_to_instrument = [t for t in node.targets if self._is_target_in_scope(t)] # First-def Name targets: not yet in scope AND not excluded @@ -588,8 +2273,9 @@ def visit_Assign(self, node: ast.Assign) -> ast.stmt | list[ast.stmt]: and not self.session_data.scope_manager.is_skip_reference_taking(t.id) ] # First-def TUPLE-unpack elements: ``a, _, _ = obj.attr`` + # (the ``[a, _, _] = obj.attr`` List spelling unpacks identically) for t in node.targets: - if isinstance(t, ast.Tuple) and not self._is_target_in_scope(t): + if isinstance(t, (ast.Tuple, ast.List)) and not self._is_target_in_scope(t): for elt in t.elts: if ( isinstance(elt, ast.Name) @@ -637,118 +2323,118 @@ def visit_Assign(self, node: ast.Assign) -> ast.stmt | list[ast.stmt]: and not self._is_meta_primitive_literal(node.value) ): first_def_targets.append(t) - # For tuple targets, snapshot which elements are in scope NOW, - # before _visit_target adds first-definitions to scope. - # Use a dict keyed by id(target) so each tuple target gets its - # own snapshot — this is needed for multi-target assignments - # like ``(a, b) = (c, d) = expr``. - tuple_scope_snapshots: dict[int, set[int]] = {} + # For tuple targets, snapshot which elements are in scope NOW, before + # _visit_target adds first-definitions to scope. Attached to the + # target itself so it survives a statement rebuild that hands back a + # different node object. for t in targets_to_instrument: - if isinstance(t, ast.Tuple): - tuple_scope_snapshots[id(t)] = { + if isinstance(t, (ast.Tuple, ast.List)): + t._pyir_scope_snapshot = { # type: ignore[union-attr] i for i, elt in enumerate(t.elts) if self._is_element_in_scope(elt) } - - for target in node.targets: - self._visit_target(target) - self.generic_visit(node) - - tuple_targets = [t for t in targets_to_instrument if isinstance(t, ast.Tuple)] - subscript_targets = [ - t for t in targets_to_instrument if isinstance(t, ast.Subscript) - ] - scalar_targets = [ - t - for t in targets_to_instrument - if not isinstance(t, (ast.Tuple, ast.Subscript)) - ] - - non_empty = sum( - 1 for group in (tuple_targets, subscript_targets, scalar_targets) if group - ) - if non_empty > 1: - raise DSLUserCodeError( - DiagId.UNSUP_MIXED_ASSIGN_TARGETS, - filename=self.session_data.file_name, - ) - - if tuple_targets: - # Evaluate the RHS into a shared temp once, then decompose - # each tuple target using its own per-target scope snapshot. - all_stmts: list[ast.stmt] = [] - shared_tmp: str | None = None - for tt in tuple_targets: - snapshot = tuple_scope_snapshots[id(tt)] - result = self._decompose_tuple_assign( - node, tt, snapshot, rhs_temp_name=shared_tmp - ) - if isinstance(result, list): - all_stmts.extend(result) - else: - # Starred or no instrumented elements — keep original - all_stmts.append(result) - # After the first target creates the temp, reuse it. - if shared_tmp is None and isinstance(result, list): - # The temp name is _pyir_tmp_N; find it from the - # statements emitted by _decompose_tuple_assign. - for s in result: - if ( - isinstance(s, ast.Assign) - and len(s.targets) == 1 - and isinstance(s.targets[0], ast.Name) - and s.targets[0].id.startswith("_pyir_tmp_") - ): - shared_tmp = s.targets[0].id - break - return all_stmts if all_stmts else node - - if subscript_targets: - exclude = {self._target_to_path_str(t) for t in subscript_targets} - # Same ordering invariant as scalar_targets above: collect RHS - # reads first so the subscript-assign sees rewritten reads. - rhs_call_func_ids = self._collect_call_func_ids(node.value) - self_reads = self._collect_attr_reads( - node.value, exclude, rhs_call_func_ids - ) - other_reads = self._collect_other_attr_reads( - node.value, exclude, rhs_call_func_ids - ) - prologue_reads: list[ast.stmt] | None = None - if self_reads or other_reads: - attr_read_result = self._insert_pyir_attr_reads(node, exclude) - if isinstance(attr_read_result, list): - prologue_reads = attr_read_result - sub_result = self._insert_pyir_subscript_assign(node, subscript_targets[0]) - if isinstance(sub_result, list) and prologue_reads is not None: - return prologue_reads[:-1] + sub_result - return sub_result - - if scalar_targets: - exclude = {self._target_to_path_str(t) for t in scalar_targets} - # Run attr-read instrumentation FIRST so node.value's reads are - # rewritten to _pyir_attr_N locals BEFORE _insert_pyir_assign - # deepcopies node into its hasattr branches. Without this, both - # ``self._x = self._x + 1`` (covered by ``_collect_attr_reads``) - # and ``obj = Cls(obj.x + 1)`` rebinding patterns (covered by - # ``_collect_other_attr_reads``) keep the raw attribute access - # in the RHS and bypass the read. - rhs_call_func_ids = self._collect_call_func_ids(node.value) - self_reads = self._collect_attr_reads( - node.value, exclude, rhs_call_func_ids - ) - other_reads = self._collect_other_attr_reads( - node.value, exclude, rhs_call_func_ids - ) - prologue_reads = None - if self_reads or other_reads: - attr_read_result = self._insert_pyir_attr_reads(node, exclude) - if isinstance(attr_read_result, list): - prologue_reads = attr_read_result - result = self._insert_pyir_assign(node, scalar_targets) - if isinstance(result, list) and prologue_reads is not None: - return prologue_reads[:-1] + result - return result - + node._pyir_first_def_targets = first_def_targets # type: ignore[attr-defined] + return targets_to_instrument + + def _assign_sequence_target(self, node: ast.Assign) -> "ast.stmt | list[ast.stmt]": + """Rebuild an assignment whose single target is a Tuple/List.""" + assert len(node.targets) == 1, "chained assign was not normalised" + target = node.targets[0] + assert isinstance(target, (ast.Tuple, ast.List)) + # Read the snapshot off the target BEFORE the rewrite below, which may + # hand back a different statement object: the fact is attached to THIS + # target, and a rebound ``node.targets`` would miss it (a lost snapshot + # silently downgrades in-scope elements to first-defs). + all_stmts: list[ast.stmt] = [] + + exclude = {self._target_to_path_str(elt) for elt in target.elts} + attr_read_result = self._insert_pyir_attr_reads(node, exclude) + if isinstance(attr_read_result, list): + all_stmts.extend(attr_read_result[:-1]) + # The reads substitute into the statement the call returns, which + # is what the decompose must pin. + rewritten = attr_read_result[-1] + assert isinstance(rewritten, ast.Assign) + node = rewritten + + snapshot: set[int] = getattr(target, "_pyir_scope_snapshot", set()) + decomposed = self._decompose_unpack_assign(node, target, snapshot) + if isinstance(decomposed, list): + all_stmts.extend(decomposed) + else: + # Starred or no instrumented elements -- keep the original node. + all_stmts.append(decomposed) + return all_stmts if all_stmts else node + + def _assign_subscript_target(self, node: ast.Assign) -> "ast.stmt | list[ast.stmt]": + """Instrument a subscript target (``c[k] = rhs``).""" + assert len(node.targets) == 1, "chained assign was not normalised" + target = node.targets[0] + assert isinstance(target, ast.Subscript) + exclude = {self._target_to_path_str(target)} + # Same ordering invariant as scalar_targets above: collect RHS + # reads first so the subscript-assign sees rewritten reads. + self_reads = self._collect_attr_reads(node.value, exclude) + other_reads = self._collect_other_attr_reads(node.value, exclude) + deep_reads = self._collect_deep_attr_reads(node.value, exclude) + global_reads = self._collect_global_name_reads(node.value, exclude) + prologue_reads: list[ast.stmt] | None = None + if self_reads or other_reads or deep_reads or global_reads: + attr_read_result = self._insert_pyir_attr_reads(node, exclude) + if isinstance(attr_read_result, list): + prologue_reads = attr_read_result + else: + # Function-scope: no read instrumentation fires here, so add + # only the superseded-generation probes (no rewrites). + probe_stmts = self._build_generation_probe_stmts(node, exclude) + if probe_stmts: + prologue_reads = [*probe_stmts, node] + sub_result = self._insert_pyir_subscript_assign(node, target) + if isinstance(sub_result, list) and prologue_reads is not None: + return prologue_reads[:-1] + sub_result + return sub_result + + def _assign_scalar_target(self, node: ast.Assign) -> "ast.stmt | list[ast.stmt]": + """Instrument an in-scope plain-name / attribute target.""" + assert len(node.targets) == 1, "chained assign was not normalised" + target = node.targets[0] + assert isinstance(target, (ast.Name, ast.Attribute)) + exclude = {self._target_to_path_str(target)} + # Run attr-read instrumentation FIRST so node.value's reads are + # rewritten to _pyir_attr_N locals BEFORE _insert_pyir_assign + # deepcopies node into its hasattr branches. Without this, both + # ``self._x = self._x + 1`` (covered by ``_collect_attr_reads``) + # and ``obj = Cls(obj.x + 1)`` rebinding patterns (covered by + # ``_collect_other_attr_reads``) keep the raw attribute access + # in the RHS and bypass the read. + self_reads = self._collect_attr_reads(node.value, exclude) + other_reads = self._collect_other_attr_reads(node.value, exclude) + deep_reads = self._collect_deep_attr_reads(node.value, exclude) + global_reads = self._collect_global_name_reads(node.value, exclude) + prologue_reads = None + if self_reads or other_reads or deep_reads or global_reads: + attr_read_result = self._insert_pyir_attr_reads(node, exclude) + if isinstance(attr_read_result, list): + prologue_reads = attr_read_result + else: + # Function-scope: no read instrumentation fires here, so add + # only the superseded-generation probes (no rewrites). + probe_stmts = self._build_generation_probe_stmts(node, exclude) + if probe_stmts: + prologue_reads = [*probe_stmts, node] + assign_result = self._insert_pyir_assign(node, [target]) + if isinstance(assign_result, list) and prologue_reads is not None: + return prologue_reads[:-1] + assign_result + return assign_result + + def _assign_first_def_fallthrough( + self, node: ast.Assign + ) -> "ast.stmt | list[ast.stmt]": + """No target needs reassignment instrumentation: still rewrite the RHS + reads, and commit any first-def target (``_collect_assign_facts`` + attached the list) so a first definition inside staged CF gets its + eager ref.""" + assert len(node.targets) == 1, "chained assign was not normalised" + first_def_targets = getattr(node, "_pyir_first_def_targets", []) # First-time definitions: still scan RHS for self.X reads, # and emit pyir_assign(name, None, name) for first-def Name # targets so that pyir_assign can create an eager ref when the @@ -764,8 +2450,105 @@ def visit_Assign(self, node: ast.Assign) -> ast.stmt | list[ast.stmt]: return [base] + first_def_stmts return base + def _split_multiple_targets(self, node: ast.Assign) -> "list[ast.stmt]": + """Rewrite ``t1 = t2 = rhs`` into one RHS temp plus one store per target. + + Python evaluates the RHS ONCE and then stores into each target left to + right, which is exactly this shape:: + + _pyir_rhs_N = rhs + t1 = _pyir_rhs_N + t2 = _pyir_rhs_N + + Every downstream exit then sees a SINGLE target, so none of them has to + carry its own chained-assign machinery -- three different ones existed + (the tuple decompose's shared pin, ``_insert_pyir_assign``'s + ``_pyir_chain_N``, and the subscript exit, which simply ignored every + target after the first). + + Only for a chained assign (``_is_multi_targets_assign``): rewriting a + single-target statement would add a temp to nearly every statement in + the program and buy nothing. + """ + tmp = f"_pyir_rhs_{self.session_data.counter}" + self.session_data.counter += 1 + # The synthetic temp is machinery-owned: no first-def choke, no ledger + # row (same treatment the nested-unpack temps get). + self.session_data.scope_manager.add_skip_reference_taking(tmp) + hoist = ast.Assign( + targets=[ast.Name(id=tmp, ctx=ast.Store())], + value=node.value, + ) + rewritten: list[ast.stmt] = [ + ast.copy_location(ast.fix_missing_locations(hoist), node) + ] + for target in node.targets: + store = ast.Assign( + targets=[_deepcopy_ast_root(target)], + value=ast.Name(id=tmp, ctx=ast.Load()), + ) + rewritten.append(ast.copy_location(ast.fix_missing_locations(store), node)) + result: list[ast.stmt] = [] + for stmt in rewritten: + visited = self.visit(stmt) + result.extend(visited if isinstance(visited, list) else [visited]) + return result + + def visit_Assign(self, node: ast.Assign) -> ast.stmt | list[ast.stmt]: + """Override to add PyIR instrumentation for reassignments. + + Everything here is statement level: the two rewrites that reduce a + statement to ONE target, the pre-walrus hoists, the facts the exits + share, and the dispatch to one of four exits -- one per target kind, + each a method below taking just the statement. + + The rewrites run outside-in: a chained assign is split into one store + per target FIRST, so each store then faces the nested-unpack question + on its own. Everything after them sees a single target, which is what + lets each exit assert its shape instead of grouping target lists. + """ + if self._is_multi_targets_assign(node): + return self._split_multiple_targets(node) + if self._is_nested_unpack_target(node): + return self._split_nested_unpack_target(node) + self._hoist_pre_walrus_reads(node, node.value) + self._hoist_pre_walrus_effects(node, node.value) + + targets_to_instrument = self._collect_assign_facts(node) + + for target in node.targets: + self._visit_target(target) + self.generic_visit(node) + + # A sequence target takes its exit whether or not any element needs + # instrumentation: that exit REPLACES the statement, so a target it + # declined would be dropped outright and its names never bound. + assert len(node.targets) == 1, ( + "multiple targets should have been split into multiple statements." + ) + target = node.targets[0] + if isinstance(target, (ast.Tuple, ast.List)): + return self._assign_sequence_target(node) + if targets_to_instrument: + if isinstance(target, ast.Subscript): + return self._assign_subscript_target(node) + return self._assign_scalar_target(node) + return self._assign_first_def_fallthrough(node) + + def visit_AnnAssign(self, node: ast.AnnAssign) -> "ast.stmt | list[ast.stmt]": + """An annotated assignment with a value IS an assignment: route it through the + full ``visit_Assign`` instrumentation.""" + if node.value is None or self._pyir_class_body_depth > 0: + return super().visit_AnnAssign(node) + assign = ast.copy_location( + ast.Assign(targets=[node.target], value=node.value), node + ) + ast.fix_missing_locations(assign) + return self.visit_Assign(assign) + def visit_Delete(self, node: ast.Delete) -> ast.stmt | list[ast.stmt]: - """Reject ``del`` of a LOOP-CARRIED local inside control flow. + """Reject ``del`` of a LOOP-CARRIED local inside control flow, and route + attribute targets through the one attribute-deletion funnel. A name defined BEFORE the enclosing for/while/if carries a value into the region; ``del``'ing it mid-body drops that carry (the following @@ -775,6 +2558,12 @@ def visit_Delete(self, node: ast.Delete) -> ast.stmt | list[ast.stmt]: INNERMOST region scope was defined inside the region (a trace-time temp) and deleting it is benign -> pass, constexpr-unrolled or dynamic. A top-level ``del`` and a skip-reference name (loop induction var) pass. + + ``del obj.attr`` rewrites to ``_pyir_delete_attr(obj, "attr")`` -- the + funnel the routed ``delattr`` also uses (it performs the native delete, + refuses inside dynamic staged CF, and records the unbind in meta flow). + Subscript targets stay native: watched containers choke in their own + ``__delitem__``, trace-internal ones keep Python protocol semantics. """ active = self.session_data.scope_manager.get_active_symbols() if len(active) > 1: @@ -795,10 +2584,39 @@ def visit_Delete(self, node: ast.Delete) -> ast.stmt | list[ast.stmt]: end_col_offset=getattr(node, "end_col_offset", None), var=tgt.id, ) - return node + if self._pyir_class_body_depth > 0 or not any( + isinstance(t, ast.Attribute) for t in node.targets + ): + return node + stmts: list[ast.stmt] = [] + for tgt in node.targets: + if isinstance(tgt, ast.Attribute): + new_stmt: ast.stmt = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_delete_attr", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ + self._target_as_load(tgt.value), + ast.Constant(value=tgt.attr), + ], + keywords=[], + ) + ) + else: + new_stmt = ast.Delete(targets=[tgt]) + stmts.append(ast.copy_location(ast.fix_missing_locations(new_stmt), node)) + return stmts def visit_AugAssign(self, node: ast.AugAssign) -> ast.AugAssign | list[ast.stmt]: """Override to add PyIR instrumentation for augmented assignments.""" + # Base load precedes all RHS evaluation in Python, so its anchor is + # emitted first. + base_anchor = self._hoist_pre_walrus_augassign_base(node) + self._hoist_pre_walrus_reads(node, node.value) + self._hoist_pre_walrus_effects(node, node.value) self._visit_target(node.target) self.generic_visit(node) @@ -809,23 +2627,38 @@ def visit_AugAssign(self, node: ast.AugAssign) -> ast.AugAssign | list[ast.stmt] return node target_path = self._target_to_path_str(node.target) + if isinstance(node.target, ast.Subscript) and self._is_subscript_skippable( + node.target + ): + return node + + # OTHER-attr parity with visit_Assign: an augassign RHS reading a + # non-self receiver must hoist through the same instrumented read. + # Deep-chain reads (``d["k"].val``, ``self.helper().x``) count too: + # visit_Assign's gate consults all three collectors. The hoist runs + # FIRST (it rewrites node.value in place) so the in-place-op rewrite + # below copies the hoisted RHS, never the raw attribute reads. + attr_reads = ( + self._collect_attr_reads(node.value, {target_path}) + or self._collect_other_attr_reads(node.value, {target_path}) + or self._collect_deep_attr_reads(node.value, {target_path}) + or self._collect_global_name_reads(node.value, {target_path}) + ) + read_stmts: "list[ast.stmt] | None" = None + if attr_reads: + hoisted = self._insert_pyir_attr_reads(node, {target_path}) + if isinstance(hoisted, list): + read_stmts = hoisted + if isinstance(node.target, ast.Subscript): - if self._is_subscript_skippable(node.target): - return node result = self._insert_pyir_subscript_augassign(node, node.target) - attr_reads = self._collect_attr_reads(node.value, {target_path}) - if attr_reads and isinstance(result, list): - read_stmts = self._insert_pyir_attr_reads(node, {target_path}) - if isinstance(read_stmts, list): - return read_stmts[:-1] + result + if read_stmts is not None and isinstance(result, list): + return read_stmts[:-1] + result return result - result = self._insert_pyir_augassign(node) - attr_reads = self._collect_attr_reads(node.value, {target_path}) - if attr_reads and isinstance(result, list): - read_stmts = self._insert_pyir_attr_reads(node, {target_path}) - if isinstance(read_stmts, list): - return read_stmts[:-1] + result + result = self._insert_pyir_augassign(node, base_anchor) + if read_stmts is not None and isinstance(result, list): + return read_stmts[:-1] + result return result def visit_Subscript(self, node: ast.Subscript) -> ast.expr: @@ -856,16 +2689,15 @@ def visit_Subscript(self, node: ast.Subscript) -> ast.expr: # _pyir_post_subscript_read("d[key]", d, key) call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "_pyir_post_subscript_read", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), args=[ ast.Constant(value=target_name), - deepcopy(node.value), # container expression - deepcopy(node.slice), # key expression + _deepcopy_ast_root(node.value), # container expression + _deepcopy_ast_root(node.slice), # key expression ], keywords=[], ) @@ -873,16 +2705,17 @@ def visit_Subscript(self, node: ast.Subscript) -> ast.expr: def visit_Return(self, node: ast.Return) -> ast.stmt | list[ast.stmt]: """Override to add PyIR instrumentation for attribute reads in returns.""" + self._hoist_pre_walrus_effects(node, node.value) self.generic_visit(node) if node.value is not None: return self._insert_pyir_attr_reads(node) return node - # Method names whose call on a meta ``list`` / ``dict`` / ``set`` - # mutates the container in place. When the call appears as a - # statement inside staged CF, ``_pyir_check_no_complex_m2m_call`` - # raises a ``DSLUserCodeError`` so users see a clear diagnostic - # instead of silent single-iteration baking. + # Method names whose call on a meta ``list`` / ``dict`` / ``set`` / + # ``collections.deque`` mutates the container in place. When the call + # appears as a statement inside staged CF, + # ``_pyir_check_no_complex_m2m_call`` raises a ``DSLUserCodeError`` so + # users see a clear diagnostic instead of silent single-iteration baking. _CONTAINER_MUTATOR_METHODS: frozenset[str] = frozenset( { "append", @@ -901,72 +2734,164 @@ def visit_Return(self, node: ast.Return) -> ast.stmt | list[ast.stmt]: "intersection_update", "difference_update", "symmetric_difference_update", + "appendleft", + "popleft", + "extendleft", + "rotate", } ) - def visit_Expr(self, node: ast.Expr) -> ast.stmt | list[ast.stmt]: - """Override to add PyIR instrumentation for attribute reads in - expressions, plus a runtime mutation guard for ``a.append(...)`` - style calls on meta Python containers inside staged CF. - """ - # Build the mutation guard BEFORE generic_visit: generic_visit rewrites - # a subscript receiver (``d["xs"]``) into a ``_pyir_post_subscript_read`` - # Call, hiding the ``.`` pattern the guard matches on. - guard = self._build_container_mutator_guard(node) - self.generic_visit(node) - result = self._insert_pyir_attr_reads(node) - if guard is None: - return result - if isinstance(result, list): - return [guard, *result] - return [guard, result] + def visit_Expr(self, node: ast.Expr) -> ast.stmt | list[ast.stmt]: + """Override to add PyIR instrumentation for attribute reads in + expressions, plus a runtime mutation guard for ``a.append(...)`` + style calls on meta Python containers inside staged CF. + """ + self._hoist_pre_walrus_effects(node, node.value) + self.generic_visit(node) + # An Attribute-rooted receiver evaluates once natively; hoist it to a + # shared temp so the guard/freeze brackets below add no re-evaluation + # (a property getter would fire once per bracketing statement). + recv_hoist = self._hoist_mutator_receiver(node) + # Build the guard from the original call BEFORE read-hoisting rewrites its + # args into later-defined temps (the guard is emitted ahead of them). + guard = self._build_container_mutator_guard(node) + # Emitted AFTER the (read-hoisted) call: it inspects the CONTAINER + # after the mutation. + freeze = self._build_container_insert_freeze(node) + result = self._insert_pyir_attr_reads(node) + stmts = result if isinstance(result, list) else [result] + if guard is not None: + stmts = [guard, *stmts] + if recv_hoist is not None: + stmts = [recv_hoist, *stmts] + if recv_hoist is None and guard is None and freeze is None: + return result + if freeze is not None: + stmts = [*stmts, freeze] + return stmts + + def _hoist_mutator_receiver(self, node: ast.Expr) -> ast.stmt | None: + """Hoist an Attribute-rooted mutator receiver into a shared temp. + + ``obj.field.append(x)`` evaluates ``obj.field`` once natively, but + the mutator guard and the insert freeze would each re-evaluate it + (a property getter fires once per bracket). Bind the receiver to + a temp ahead of the brackets and rewrite the call so the brackets + (which read ``call.func.value`` after this) share the one + evaluation. Name receivers stay put: re-reading a local is + effect-free.""" + call = node.value + if not isinstance(call, ast.Call): + return None + func = call.func + if not isinstance(func, ast.Attribute): + return None + if ( + func.attr not in self._CONTAINER_MUTATOR_METHODS + and func.attr not in self._CONTAINER_INSERT_METHODS + ): + return None + if not isinstance(func.value, ast.Attribute): + return None + recv_name = f"_pyir_recv_{self.session_data.counter}" + self.session_data.counter += 1 + # The guard's diagnostic keeps the user's spelling, not the temp. + call._pyir_recv_spelling = _unparse_safe(func.value) # type: ignore[attr-defined] + hoist = ast.Assign( + targets=[ast.Name(id=recv_name, ctx=ast.Store())], + value=func.value, + ) + func.value = ast.Name(id=recv_name, ctx=ast.Load()) + return ast.copy_location(ast.fix_missing_locations(hoist), node) + + def _build_container_mutator_guard(self, node: ast.Expr) -> ast.stmt | None: + """Return an ``_pyir_check_no_complex_m2m_call(container, ...)`` + statement to prepend before *node* when *node* is a statement- + level call of the form ``.(...)``. Returns + ``None`` when the pattern doesn't match. + """ + call = node.value + if not isinstance(call, ast.Call): + return None + func = call.func + if not isinstance(func, ast.Attribute): + return None + if func.attr not in self._CONTAINER_MUTATOR_METHODS: + return None + container = func.value + # Simple container expressions only; an Attribute-rooted receiver + # arrives here as the shared temp `_hoist_mutator_receiver` bound. + if not isinstance(container, (ast.Name, ast.Attribute)): + return None + # Pass values only when every arg is side-effect-free to re-evaluate; + # else ``None`` (conservative). + if not call.keywords and all( + isinstance(a, (ast.Name, ast.Constant)) for a in call.args + ): + values_arg: ast.expr = ast.List( + elts=[_deepcopy_ast_root(a) for a in call.args], ctx=ast.Load() + ) + else: + values_arg = ast.Constant(value=None) + return ast.copy_location( + ast.fix_missing_locations( + ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_check_no_complex_m2m_call", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ + _deepcopy_ast_root(container), + ast.Constant(value=func.attr), + ast.Constant( + value=getattr(call, "_pyir_recv_spelling", None) + or _unparse_safe(container) + ), + ast.Constant(value=self.session_data.file_name), + ast.Constant(value=node.lineno), + values_arg, + ], + keywords=[], + ), + ) + ), + node, + ) + + # Insert-like mutators add the call's value argument; a STAGED appended value + # re-resolving a loop-carried slot freezes to its current SSA (see the freeze). + _CONTAINER_INSERT_METHODS: frozenset[str] = frozenset( + {"append", "extend", "insert", "update", "setdefault"} + ) - def _build_container_mutator_guard(self, node: ast.Expr) -> ast.stmt | None: - """Return an ``_pyir_check_no_complex_m2m_call(container, ...)`` - statement to prepend before *node* when *node* is a statement- - level call of the form ``.(...)``. Returns - ``None`` when the pattern doesn't match. - """ + def _build_container_insert_freeze(self, node: ast.Expr) -> ast.stmt | None: + """Emit ``_pyir_freeze_staged_container_inserts(container, method)`` after a + statement-level insert-mutator call; gating lives in the runtime.""" call = node.value if not isinstance(call, ast.Call): return None func = call.func if not isinstance(func, ast.Attribute): return None - if func.attr not in self._CONTAINER_MUTATOR_METHODS: + if func.attr not in self._CONTAINER_INSERT_METHODS: return None container = func.value - # Only instrument simple container expressions; complex - # container expressions (e.g. ``obj.field.list``) would require - # a temporary to avoid re-evaluation and are out of scope for - # this guard. - if isinstance(container, (ast.Name, ast.Attribute)): - pass - elif isinstance(container, ast.Subscript) and not any( - isinstance(_n, ast.Call) for _n in ast.walk(container) - ): - # A dict/list-held container reached one indirection deep - # (``d["xs"].append``); accept only side-effect-free receivers so - # re-evaluating the receiver for the guard cannot double-execute. - pass - else: + if not isinstance(container, (ast.Name, ast.Attribute)): return None return ast.copy_location( ast.fix_missing_locations( ast.Expr( value=ast.Call( - func=_create_module_attribute( - "_pyir_check_no_complex_m2m_call", - submodule_name="pyir_runtime", + func=_create_runtime_attribute( + "_pyir_freeze_staged_container_inserts", lineno=node.lineno, col_offset=node.col_offset, ), args=[ - deepcopy(container), + _deepcopy_ast_root(container), ast.Constant(value=func.attr), - ast.Constant(value=_unparse_safe(container)), - ast.Constant(value=self.session_data.file_name), - ast.Constant(value=node.lineno), ], keywords=[], ), @@ -977,10 +2902,14 @@ def _build_container_mutator_guard(self, node: ast.Expr) -> ast.stmt | None: def _handle_constexpr_while(self, node: ast.While) -> list[ast.stmt]: """Override to add PyIR scope isolation for const_expr while loops.""" - # Visit test expression outside branch scopes. + # Visit test expression outside branch scopes; test-position reads + # skip the statement hoist, so they take record-only chokes here. self.visit(node.test) + node.test = self._wrap_test_observation_reads(node.test) # Visit body in its own scope so first-definitions don't leak. - node.body = self._visit_stmts_in_cf_scope(node.body) + born: set[str] = set() + node.body = self._visit_stmts_in_cf_scope(node.body, collect_bindings=born) + self._readd_constexpr_arm_bindings(born) # Bracket the (trace-time-unrolled) body so mutations directly # inside it are treated as constexpr-governed by the guards. node.body = self._wrap_body_in_constexpr_scope(node, node.body) @@ -990,24 +2919,17 @@ def _handle_constexpr_while(self, node: ast.While) -> list[ast.stmt]: def _handle_constexpr_if(self, node: ast.If) -> list[ast.stmt]: """Override to add PyIR scope isolation for const_expr if statements.""" - # Visit test expression outside branch scopes. - has_else = bool(node.orelse) + # Visit test expression outside branch scopes; test-position reads + # skip the statement hoist, so they take record-only chokes here. self.visit(node.test) + node.test = self._wrap_test_observation_reads(node.test) # Visit each branch in its own scope so first-definitions in one # branch don't leak into sibling branches (fixes UnboundLocalError # when only one const_expr branch runs at runtime). - node.body, body_definitions = ( - self._visit_stmts_in_cf_scope_and_collect_definitions(node.body) - ) - node.orelse, else_definitions = ( - self._visit_stmts_in_cf_scope_and_collect_definitions(node.orelse) - ) - # An exhaustive if/else definitely initializes names defined by every - # arm. Propagate only that intersection; one-sided definitions remain - # out of scope and retain Python's possible-UnboundLocal semantics. - if has_else: - for name in body_definitions & else_definitions: - self.session_data.scope_manager.add_to_scope(name) + born: set[str] = set() + node.body = self._visit_stmts_in_cf_scope(node.body, collect_bindings=born) + node.orelse = self._visit_stmts_in_cf_scope(node.orelse, collect_bindings=born) + self._readd_constexpr_arm_bindings(born) # Bracket each branch (the trace-time-selected one runs once) so # mutations directly inside it are treated as constexpr-governed. node.body = self._wrap_body_in_constexpr_scope(node, node.body) @@ -1020,21 +2942,18 @@ def _handle_constexpr_elif(self, elif_node: ast.If) -> ast.stmt: """Override to add PyIR scope isolation for const_expr elif nodes.""" # Visit test outside branch scopes; visit each # branch in its own scope to prevent cross-branch - # first-definition leakage. - has_else = bool(elif_node.orelse) + # first-definition leakage. Test-position reads skip the statement + # hoist, so they take record-only chokes here. self.visit(elif_node.test) - elif_node.body, body_definitions = ( - self._visit_stmts_in_cf_scope_and_collect_definitions(elif_node.body) - ) - elif_node.orelse, else_definitions = ( - self._visit_stmts_in_cf_scope_and_collect_definitions(elif_node.orelse) - ) - # Feed definitions from a complete elif/else tail into its parent's - # temporary branch scope. The outer if then intersects that set with - # its own body, handling arbitrary exhaustive elif chains. - if has_else: - for name in body_definitions & else_definitions: - self.session_data.scope_manager.add_to_scope(name) + elif_node.test = self._wrap_test_observation_reads(elif_node.test) + born: set[str] = set() + elif_node.body = self._visit_stmts_in_cf_scope( + elif_node.body, collect_bindings=born + ) + elif_node.orelse = self._visit_stmts_in_cf_scope( + elif_node.orelse, collect_bindings=born + ) + self._readd_constexpr_arm_bindings(born) # Bracket each branch so mutations directly inside the # trace-time-selected branch are treated as constexpr-governed. elif_node.body = self._wrap_body_in_constexpr_scope(elif_node, elif_node.body) @@ -1044,139 +2963,324 @@ def _handle_constexpr_elif(self, elif_node: ast.If) -> ast.stmt: assert isinstance(elif_node.test, ast.Call) return self._insert_cf_symbol_check(elif_node.test.func) - def visit_While(self, node: ast.While) -> "ast.While | list[ast.stmt]": - """Route a staged ``while`` condition through ``_pyir_while_cond``. + def _wrap_test_observation_reads(self, expr: ast.expr) -> ast.expr: + """Record-only observation chokes for symbol reads in TEST position: + test expressions never pass through the statement-level read hoist, so + their meta reads bake unrecorded. Wraps each read INLINE (same + evaluation site and count), never altering staging semantics.""" + mod_name = self._reading_module_name() + fn_globals = self.session_data.function_globals or {} + active = self.session_data.scope_manager.get_active_symbols() + bound = _locally_bound_names(expr) + sm = self.session_data.scope_manager + outer = self - The base preprocessor moves ``node.test`` into ``while_before_block`` - verbatim, where a watched fold bakes ``scf.condition(true)`` -- an - unkillable runtime hang once the body mutates a condition slot. - Wrapping the condition in ``_pyir_while_cond`` BEFORE ``super()`` - embeds the lift inside the before-block (evaluated per iteration), - so a to-be-stored slot lifts to a dynamic ``scf.condition(cmpi...)``. + def in_scope(name: str) -> bool: + return any(name in scope for scope in active) - Constexpr whiles keep the base behavior (they unroll at trace time). - """ - if self.is_node_constexpr(node): - return super().visit_While(node) - node.test = ast.copy_location( - ast.Call( - func=_create_module_attribute( - "_pyir_while_cond", - submodule_name="pyir_runtime", - lineno=getattr(node.test, "lineno", None), - col_offset=getattr(node.test, "col_offset", None), - ), - args=[node.test], - keywords=[], - ), - node.test, - ) - return super().visit_While(node) + class _Wrapper(ast.NodeTransformer): + # Late-bound bodies keep their raw reads (deferred contract). + def visit_Lambda(self, n: ast.Lambda) -> ast.Lambda: + return n + + def visit_FunctionDef(self, n: ast.FunctionDef) -> ast.FunctionDef: + return n + + def visit_Call(self, n: ast.Call) -> ast.Call: + # The call target is never a value read; args/keywords are. + n.args = [self.visit(a) for a in n.args] + for kw in n.keywords: + kw.value = self.visit(kw.value) + return n + + def visit_Attribute(self, n: ast.Attribute) -> ast.expr: + if not isinstance(n.ctx, ast.Load): + return n + root: ast.expr = n.value + while isinstance(root, ast.Attribute): + root = root.value + if isinstance(root, ast.Name) and root.id in ( + "__base_dsl__", + "__module_dsl__", + _PYIR_RUNTIME_ALIAS, + ): + return n + n.value = self.visit(n.value) + spelling = outer._access_path_str(n) or f".{n.attr}" + call = ast.Call( + func=_create_runtime_attribute("_pyir_obs_read"), + args=[ + ast.Constant(value=spelling), + n.value, + ast.Constant(value=n.attr), + ast.Constant(value=mod_name), + ], + keywords=[], + ) + return ast.copy_location(ast.fix_missing_locations(call), n) - def _prepare_while_condition_vars( + def visit_Name(self, n: ast.Name) -> ast.expr: + if not isinstance(n.ctx, ast.Load): + return n + name = n.id + if ( + name in bound + or name.startswith("_pyir_") + or name in ("__base_dsl__", "__module_dsl__", _PYIR_RUNTIME_ALIAS) + or sm.is_skip_reference_taking(name) + or in_scope(name) + or name not in fn_globals + ): + return n + obj = fn_globals.get(name) + if isinstance( + obj, + ( + types.ModuleType, + type, + types.FunctionType, + types.BuiltinFunctionType, + ), + ): + return n + call = ast.Call( + func=_create_runtime_attribute("_pyir_obs_global_read"), + args=[ + ast.Constant(value=name), + ast.Name(id=name, ctx=ast.Load()), + ast.Constant(value=mod_name), + ], + keywords=[], + ) + return ast.copy_location(ast.fix_missing_locations(call), n) + + wrapped = _Wrapper().visit(expr) + ast.fix_missing_locations(wrapped) + return wrapped + + def create_if_function( self, - node: ast.While, + func_name: str, + node: ast.If, write_args: list[str], - while_before_stmts: list[ast.stmt], - ) -> list[ast.stmt]: - """Override to insert PyIR-specific pyir_read for write_args in condition. + full_write_args_count: int, + ) -> ast.FunctionDef: + """Override: observe symbol reads in the dynamic-if TEST before the + region synthesis embeds it (tests skip the statement-level hoist).""" + node.test = self._wrap_test_observation_reads(node.test) + return super().create_if_function( + func_name, node, write_args, full_write_args_count + ) - Without this, literal-backed values (e.g. Int32(0)) rematerialize - arith.constant instead of loading from the pyir.ref. + @staticmethod + def _attr_chain_root_path(node: ast.expr) -> "tuple[str, tuple[str, ...]] | None": + """``(root_name, hops)`` for a pure Name-rooted attribute chain, e.g. + ``c0.inner.v0`` -> ``("c0", ("inner", "v0"))``; ``None`` otherwise.""" + hops: list[str] = [] + cur = node + while isinstance(cur, ast.Attribute): + hops.append(cur.attr) + cur = cur.value + if isinstance(cur, ast.Name) and hops: + return cur.id, tuple(reversed(hops)) + return None - Also prepends the ``pyir_tag_pending_writes`` write-set prologue so - ``_pyir_while_cond`` can see the body's not-yet-landed stores when it - decides fold-vs-lift. Both land at the top of - ``while_before_block`` (evaluated once per iteration, before the - condition): the tag FIRST so the write-set is populated before the - operand reload and the condition eval. - """ - prep = self._build_pyir_read_prologue(node, write_args, helper="pyir_read") - tag_stmt = self._build_pending_write_tags(node, write_args) - return tag_stmt + prep + @classmethod + def _while_cond_carried_attr_paths( + cls, node: ast.While + ) -> "dict[str, tuple[tuple[str, ...], ...]]": + """Attribute paths both READ in the while condition and WRITTEN in the + body, per root name (the runtime promotes and re-loads exactly these).""" + read_paths: "dict[str, set[tuple[str, ...]]]" = {} + for n in ast.walk(node.test): + if isinstance(n, ast.Attribute) and isinstance(n.ctx, ast.Load): + chain = cls._attr_chain_root_path(n) + if chain is not None: + read_paths.setdefault(chain[0], set()).add(chain[1]) + # Keep only maximal chains: ast.walk yields every sub-chain of a + # dotted read; a proper prefix is the traversal, not the read. + for root, paths in read_paths.items(): + read_paths[root] = { + p for p in paths if not any(q != p and q[: len(p)] == p for q in paths) + } + + write_paths: "dict[str, set[tuple[str, ...]]]" = {} + + class _BodyWriteWalker(ast.NodeVisitor): + def _record(self, tgt: ast.expr) -> None: + if isinstance(tgt, ast.Attribute): + chain = cls._attr_chain_root_path(tgt) + if chain is not None: + write_paths.setdefault(chain[0], set()).add(chain[1]) + elif isinstance(tgt, ast.Name): + # A whole-name rebind replaces the object, so every + # condition-read attr path rooted at it is written. + write_paths.setdefault(tgt.id, set()).update( + read_paths.get(tgt.id, ()) + ) + elif isinstance(tgt, (ast.Tuple, ast.List)): + for e in tgt.elts: + self._record(e) + elif isinstance(tgt, ast.Starred): + self._record(tgt.value) + + def visit_Assign(self, n: ast.Assign) -> None: + for t in n.targets: + self._record(t) + self.generic_visit(n) + + def visit_AugAssign(self, n: ast.AugAssign) -> None: + self._record(n.target) + self.generic_visit(n) + + def visit_AnnAssign(self, n: ast.AnnAssign) -> None: + self._record(n.target) + self.generic_visit(n) + + # Nested scopes analyze themselves; their writes are not this + # loop body's writes. + def visit_FunctionDef(self, n: ast.FunctionDef) -> None: + return + + def visit_AsyncFunctionDef(self, n: ast.AsyncFunctionDef) -> None: + return + + def visit_Lambda(self, n: ast.Lambda) -> None: + return + + def visit_ClassDef(self, n: ast.ClassDef) -> None: + return + + walker = _BodyWriteWalker() + for stmt in node.body: + walker.visit(stmt) + + out: "dict[str, tuple[tuple[str, ...], ...]]" = {} + for root, rpaths in read_paths.items(): + carried = rpaths & write_paths.get(root, set()) + if carried: + out[root] = tuple(sorted(carried)) + return out - def _build_pending_write_tags( + def _prepare_while_condition_vars( self, node: ast.While, write_args: list[str], + while_before_stmts: list[ast.stmt], ) -> list[ast.stmt]: - """Build the ``pyir_tag_pending_writes(...)`` before-block prologue. - - Collects the while BODY's syntactic write-set: the plain-name - write_args, plus attribute/subscript targets whose owner Name also - appears in the CONDITION (only those can affect the condition's - fold, and only those are guaranteed readable in the before-block - scope). See ``pyir_tag_pending_writes`` for why this must run - before the condition evaluates. - """ - cond_names = { - n.id - for n in ast.walk(node.test) - if isinstance(n, ast.Name) and isinstance(n.ctx, ast.Load) - } - specs: list[ast.expr] = [] - seen: set[str] = set() - for name in write_args: - if name in seen: - continue - seen.add(name) - specs.append( - ast.Tuple( - elts=[ - ast.Constant(value=name), - ast.Constant(value=None), - ast.Constant(value=None), - ], - ctx=ast.Load(), - ) + """Insert a PyIR write_arg prologue in the ``scf.while`` condition so + carried values load their cells before the condition reads them.""" + # is_condition_read=False is exactness-gated: emitted only when the + # condition's read set is fully AST-enumerable, else promote-by-default. + opaque_test = any( + isinstance( + sub, + ( + ast.Call, + ast.Lambda, + ast.GeneratorExp, + ast.ListComp, + ast.SetComp, + ast.DictComp, + ast.Subscript, + ast.Starred, + ast.Await, + ), ) - for stmt in ast.walk(node): - targets: list[ast.expr] = [] - if isinstance(stmt, ast.Assign): - targets = list(stmt.targets) - elif isinstance(stmt, (ast.AugAssign, ast.AnnAssign)): - targets = [stmt.target] - for t in targets: - owner_expr: ast.expr | None = None - slot_const: ast.expr | None = None - if isinstance(t, ast.Attribute) and isinstance(t.value, ast.Name): - owner_expr = ast.Name(id=t.value.id, ctx=ast.Load()) - slot_const = ast.Constant(value=t.attr) - dotted = f"{t.value.id}.{t.attr}" - elif ( - isinstance(t, ast.Subscript) - and isinstance(t.value, ast.Name) - and isinstance(t.slice, ast.Constant) - ): - owner_expr = ast.Name(id=t.value.id, ctx=ast.Load()) - slot_const = ast.Constant(value=t.slice.value) - dotted = f"{t.value.id}[{t.slice.value!r}]" - else: - continue - if t.value.id not in cond_names or dotted in seen: - continue - seen.add(dotted) - specs.append( + for sub in ast.walk(node.test) + ) + condition_read: "set[str] | None" + if opaque_test: + condition_read = None + else: + condition_read = { + name.id + for name in ast.walk(node.test) + if isinstance(name, ast.Name) and isinstance(name.ctx, ast.Load) + } + return self._build_pyir_read_prologue( + node, + write_args, + helper="pyir_promote_while_carried_arg", + condition_read=condition_read, + attr_paths=self._while_cond_carried_attr_paths(node), + attr_paths_kwarg="condition_attr_paths", + ) + + @staticmethod + def _region_facts_tag(node: "ast.For | ast.While") -> "ast.Call | None": + """The ``pyir_tag_region_attr_writes`` decorator call carrying the + declared write facts of *node*'s ORIGINAL body: the ``(base, attr)`` + pairs it assigns directly, the ``(base, method)`` pairs it calls on + Name-rooted receiver paths, and the ``(func_name, arg_base)`` pairs it + calls as free functions with a Name-rooted first argument; ``None`` + with no facts.""" + attr_writes = sorted(_collect_direct_attr_write_pairs(node.body)) + method_calls = sorted(_collect_receiver_method_call_pairs(node.body)) + free_calls = sorted(_collect_free_call_arg_pairs(node.body)) + if not (attr_writes or method_calls or free_calls): + return None + + def _pairs_tuple(pairs: "list[tuple[str, str]]") -> "ast.Tuple": + return ast.Tuple( + elts=[ ast.Tuple( - elts=[ast.Constant(value=dotted), owner_expr, slot_const], + elts=[ + ast.Constant(value=b), + ast.Constant(value=a), + ], ctx=ast.Load(), ) - ) - if not specs: - return [] - call = ast.Expr( - value=ast.Call( - func=_create_module_attribute( - "pyir_tag_pending_writes", - submodule_name="pyir_runtime", - lineno=node.lineno, - col_offset=node.col_offset, - ), - args=specs, - keywords=[], + for b, a in pairs + ], + ctx=ast.Load(), ) + + tag = ast.Call( + func=_create_runtime_attribute( + "pyir_tag_region_attr_writes", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ + _pairs_tuple(attr_writes), + _pairs_tuple(method_calls), + _pairs_tuple(free_calls), + ], + keywords=[], ) - return [ast.copy_location(ast.fix_missing_locations(call), node)] + ast.copy_location(tag, node) + ast.fix_missing_locations(tag) + return tag + + @override + def create_loop_function( + self, func_name: str, node: ast.For, *args: Any, **kwargs: Any + ) -> ast.FunctionDef: + """Tag the staged loop-body function with the declared write facts of + its ORIGINAL body (direct assigns, receiver method calls, free calls).""" + tag = self._region_facts_tag(node) + func_def = super().create_loop_function(func_name, node, *args, **kwargs) + if tag is not None: + func_def.decorator_list.append(tag) + return func_def + + @override + def create_while_function( + self, func_name: str, node: ast.While, *args: Any, **kwargs: Any + ) -> ast.FunctionDef: + """Tag the staged while AFTER block (the loop body) with the declared + write facts of its ORIGINAL body, same vocabulary as the for tag.""" + tag = self._region_facts_tag(node) + func_def = super().create_while_function(func_name, node, *args, **kwargs) + if tag is not None: + for stmt in func_def.body: + if isinstance(stmt, ast.FunctionDef) and stmt.name.startswith( + "while_after_block_" + ): + stmt.decorator_list.append(tag) + break + return func_def def _prepare_loop_body_vars( self, @@ -1191,16 +3295,171 @@ def _prepare_loop_body_vars( instead of the loop-carried value. """ needs_prologue = [ - var for var in write_args if self._var_read_before_write(node.body, var) + var + for var in write_args + if self._var_read_before_write(node.body, var) + # A var whose precise first reference is a pure (re)definition is a + # per-iteration recompute, not a carry; promoting it would demote it. + and self._classify_first_ref(node.body, var) != "store" ] - return self._build_pyir_read_prologue(node, needs_prologue) + # ``bare_first_use``: True when the first body reference is a plain-use + # Load (no ref created), False for an assignment-context read. + bare_first_use = { + var: self._var_first_ref_is_bare_use(node.body, var) + for var in needs_prologue + } + # The synthetic live-out carry is write-only in the body, so the scan + # above cannot flag it: force it into the prologue as a bare first use. + liveout_carry = self.session_data.liveout_loop_carried_vars.get(id(node)) + if ( + liveout_carry is not None + and liveout_carry in write_args + and liveout_carry not in bare_first_use + ): + needs_prologue.append(liveout_carry) + bare_first_use[liveout_carry] = True + # Re-load the same declared carried attr paths at the after-block top: + # a cross-region use of the before-block SSA would break dominance. + attr_paths: "dict[str, tuple[tuple[str, ...], ...]]" = {} + if isinstance(node, ast.While): + attr_paths = self._while_cond_carried_attr_paths(node) + for root in attr_paths: + if root in write_args and root not in needs_prologue: + needs_prologue.append(root) + bare_first_use[root] = self._var_first_ref_is_bare_use( + node.body, root + ) + # WHOLE-Name-rebound fact (list carry gate): a rebound name marks a genuine + # loop-carried list; element-only writes keep the watched-choke lifecycle. + whole_rebound = { + var: var in _whole_name_rebound_names(node.body) for var in needs_prologue + } + return self._build_pyir_read_prologue( + node, + needs_prologue, + bare_first_use=bare_first_use, + attr_paths=attr_paths, + attr_paths_kwarg="carried_attr_paths", + whole_rebound=whole_rebound, + ) + + @classmethod + def _var_first_ref_is_bare_use(cls, body: list[ast.stmt], var: str) -> bool: + """True iff the FIRST reference to *var* in *body* is a plain-use Load + (creates no ref), not the read of an assignment (which creates the ref).""" + verdict = cls._classify_first_ref(body, var) + # "assign" creates its own ref; "store" and ``None`` are not bare uses. + return verdict == "bare" + + @classmethod + def _classify_first_ref(cls, stmts: list[ast.stmt], var: str) -> "str | None": + """First reference to *var* in execution order: ``"assign"`` (read of a + self/augmented assign), ``"store"``, ``"bare"`` Load, or ``None``.""" + + def reads_var(node: "ast.AST | None") -> bool: + if node is None: + return False + for sub in ast.walk(node): + if ( + isinstance(sub, ast.Name) + and sub.id == var + and isinstance(sub.ctx, ast.Load) + ): + return True + return False + + def stores_var(target: ast.expr) -> bool: + for sub in ast.walk(target): + if ( + isinstance(sub, ast.Name) + and sub.id == var + and isinstance(sub.ctx, ast.Store) + ): + return True + return False + + def classify_block(block: "list[ast.stmt] | None") -> "str | None": + return cls._classify_first_ref(block, var) if block else None + + for stmt in stmts: + # A value that reads *var* is the assignment-context read; a pure + # store (RHS does not read *var*) (re)defines it before any read. + if isinstance(stmt, ast.Assign) and any( + stores_var(t) for t in stmt.targets + ): + if not reads_var(stmt.value): + return "store" + # A CHAINED assign is normalised into ``tmp = rhs`` plus one + # store per target, so the read lands in the hoist as a plain + # use that creates no ref -- exactly the hand-written spelling, + # which the prologue has to promote. Classifying it "assign" + # would leave nobody creating the ref before the RHS runs. + return "bare" if len(stmt.targets) > 1 else "assign" + if isinstance(stmt, ast.AnnAssign) and stores_var(stmt.target): + return "assign" if reads_var(stmt.value) else "store" + if ( + isinstance(stmt, ast.AugAssign) + and isinstance(stmt.target, ast.Name) + and stmt.target.id == var + ): + return "assign" + + # Compound statements: the guard/iterable runs first, then the + # sub-blocks in execution order. + if isinstance(stmt, ast.If): + if reads_var(stmt.test): + return "bare" + for blk in (stmt.body, stmt.orelse): + v = classify_block(blk) + if v is not None: + return v + continue + if isinstance(stmt, (ast.For, ast.AsyncFor)): + if reads_var(stmt.iter): + return "bare" + for blk in (stmt.body, stmt.orelse): + v = classify_block(blk) + if v is not None: + return v + continue + if isinstance(stmt, (ast.While,)): + if reads_var(stmt.test): + return "bare" + for blk in (stmt.body, stmt.orelse): + v = classify_block(blk) + if v is not None: + return v + continue + if isinstance(stmt, (ast.With, ast.AsyncWith)): + for item in stmt.items: + if reads_var(item.context_expr): + return "bare" + v = classify_block(stmt.body) + if v is not None: + return v + continue + + # Any other statement referencing *var*: first reference is a plain use. + if reads_var(stmt): + return "bare" + return None @staticmethod def _var_read_before_write(body: list[ast.stmt], var: str) -> bool: - """Returns True iff *var* is read (Load context) before any Store - of itself anywhere in *body*. Walks every stmt -- a Load found - before a Store wins; once a Store is encountered (Assign / - AnnAssign / AugAssign of a matching Name), the function bails.""" + """Returns True iff *var* is read before -- or within the statement that + first stores -- it (an Assign/AugAssign binds only after the RHS reads).""" + + def reads_var(node: "ast.expr | None") -> bool: + if node is None: + return False + for sub in ast.walk(node): + if ( + isinstance(sub, ast.Name) + and sub.id == var + and isinstance(sub.ctx, ast.Load) + ): + return True + return False def stores_var(target: ast.expr) -> bool: for sub in ast.walk(target): @@ -1216,15 +3475,17 @@ def stores_var(target: ast.expr) -> bool: if isinstance(stmt, ast.Assign) and any( stores_var(t) for t in stmt.targets ): - return False + # ``x = f(x)`` reads before the bind; ``x = f(other)`` does not. + return reads_var(stmt.value) if isinstance(stmt, ast.AnnAssign) and stores_var(stmt.target): - return False + return reads_var(stmt.value) if ( isinstance(stmt, ast.AugAssign) and isinstance(stmt.target, ast.Name) and stmt.target.id == var ): - return False + # ``x += expr`` always reads ``x`` first. + return True for sub in ast.walk(stmt): if ( isinstance(sub, ast.Name) @@ -1239,13 +3500,63 @@ def _build_pyir_read_prologue( node: ast.stmt, write_args: list[str], helper: str = "pyir_promote_loop_body_arg", + bare_first_use: "dict[str, bool] | None" = None, + condition_read: "set[str] | None" = None, + attr_paths: "dict[str, tuple[tuple[str, ...], ...]] | None" = None, + attr_paths_kwarg: str = "condition_attr_paths", + whole_rebound: "dict[str, bool] | None" = None, ) -> list[ast.stmt]: read_stmts: list[ast.stmt] = [] for var in write_args: + keywords: list[ast.keyword] = [] + # Only ``pyir_promote_loop_body_arg`` reads ``whole_rebound``, and + # only the True case changes behaviour (the list-carry gate). + if whole_rebound is not None and whole_rebound.get(var): + keywords.append( + ast.keyword( + arg="whole_rebound", + value=ast.Constant(value=True), + ) + ) + # The declared first-reference shape: gates the staged-scalar + # body-entry ref materialisation in the runtime hook. + if bare_first_use is not None and var in bare_first_use: + keywords.append( + ast.keyword( + arg="bare_first_use", + value=ast.Constant(value=bare_first_use[var]), + ) + ) + # Emit ``is_condition_read=False`` only for a write_arg the condition + # does not read, so a recomputed meta-primitive stays a Python value. + if condition_read is not None and var not in condition_read: + keywords.append( + ast.keyword( + arg="is_condition_read", + value=ast.Constant(value=False), + ) + ) + # Declared carried attr paths for this root, emitted as a + # tuple-of-tuples of literal hop names. + if attr_paths and attr_paths.get(var): + keywords.append( + ast.keyword( + arg=attr_paths_kwarg, + value=ast.Tuple( + elts=[ + ast.Tuple( + elts=[ast.Constant(value=hop) for hop in path], + ctx=ast.Load(), + ) + for path in attr_paths[var] + ], + ctx=ast.Load(), + ), + ) + ) pyir_read_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( helper, - submodule_name="pyir_runtime", lineno=node.lineno, col_offset=node.col_offset, ), @@ -1253,7 +3564,7 @@ def _build_pyir_read_prologue( ast.Constant(value=var), ast.Name(id=var, ctx=ast.Load()), ], - keywords=[], + keywords=keywords, ) read_stmt = ast.Assign( targets=[ast.Name(id=var, ctx=ast.Store())], @@ -1287,7 +3598,7 @@ def _slot_kwargs_for(self, target: ast.expr) -> list[ast.keyword]: ] if isinstance(target, ast.Subscript): owner_node = self._target_as_load(target.value) - slot_node = deepcopy(target.slice) + slot_node = _deepcopy_ast_root(target.slice) if hasattr(slot_node, "ctx"): slot_node.ctx = ast.Load() return [ @@ -1296,6 +3607,51 @@ def _slot_kwargs_for(self, target: ast.expr) -> list[ast.keyword]: ] return [] + @staticmethod + def _nested_fresh_container_paths(value_node: ast.expr) -> "tuple[tuple, ...]": + """Constant-keyed paths of NESTED container constructions in a fresh literal + RHS; alias entries, computed keys, and call results never appear.""" + _CTOR_NODES = ( + ast.List, + ast.Dict, + ast.Set, + ast.ListComp, + ast.SetComp, + ast.DictComp, + ) + paths: list[tuple] = [] + + def _walk(node: ast.expr, prefix: tuple) -> None: + if isinstance(node, ast.Dict): + for k_node, v_node in zip(node.keys, node.values): + # ``None`` key = ``**expansion``; non-Constant = computed key. + if k_node is None or not isinstance(k_node, ast.Constant): + continue + if isinstance(v_node, _CTOR_NODES): + p = prefix + (k_node.value,) + paths.append(p) + _walk(v_node, p) + elif isinstance(node, ast.List): + for i, v_node in enumerate(node.elts): + if isinstance(v_node, _CTOR_NODES): + p = prefix + (i,) + paths.append(p) + _walk(v_node, p) + + _walk(value_node, ()) + return tuple(paths) + + @staticmethod + def _callee_root_name(func: ast.expr) -> "str | None": + """Root ``ast.Name`` id of a callee that is a pure attribute-load chain + (``Ctor`` / ``pkg.mod.Ctor``), else ``None`` (not re-loadable as a witness).""" + node = func + while isinstance(node, ast.Attribute): + node = node.value + if isinstance(node, ast.Name): + return node.id + return None + def _target_to_path_str(self, target: ast.expr) -> str: """Convert an AST target node to a dotted path string for logging.""" if isinstance(target, ast.Name): @@ -1311,7 +3667,7 @@ def _target_to_path_str(self, target: ast.expr) -> str: def _target_as_load(self, target: ast.expr) -> ast.expr: """Deep-copy an assignment target and set context to Load.""" - t = deepcopy(target) + t = _deepcopy_ast_root(target) if isinstance(t, (ast.Name, ast.Attribute, ast.Subscript, ast.Starred)): t.ctx = ast.Load() # Recursively fix nested ctx (e.g., for a.b where a has Store ctx) @@ -1334,7 +3690,7 @@ def _subscript_safe_read(self, target: ast.expr) -> ast.expr: ): return ast.Call( func=ast.Attribute( - value=deepcopy(target.value), + value=_deepcopy_ast_root(target.value), attr="get", ctx=ast.Load(), ), @@ -1344,10 +3700,7 @@ def _subscript_safe_read(self, target: ast.expr) -> ast.expr: return self._target_as_load(target) def _collect_attr_reads( - self, - node: ast.AST, - exclude_paths: set[str] | None = None, - call_func_ids: set[int] | None = None, + self, node: ast.AST, exclude_paths: set[str] | None = None ) -> list[tuple[str, str, str]]: """Collect ``self.X`` attribute reads from *node*. @@ -1356,16 +3709,22 @@ def _collect_attr_reads( ``ast.Name`` ``self`` -- i.e. ``self.x`` but not ``self.x.y``. Skips attributes that are call targets (``self.advance()``), and any path in *exclude_paths*. - - *call_func_ids* lets a caller that already walked *node* pass the - call-target id set in, so we skip a redundant ``ast.walk``; it is - computed on demand when omitted. """ if exclude_paths is None: exclude_paths = set() - if call_func_ids is None: - call_func_ids = self._collect_call_func_ids(node) + call_func_ids = self._collect_call_func_ids(node) + # Comprehension/lambda-local names live in a private scope; must not be + # hoisted to the statement-scope prologue. + bound = _locally_bound_names(node) + # Lambda-body / genexp reads are LATE-BOUND by Python; hoisting them + # would early-bind the capture at statement time. + deferred = _deferred_execution_node_ids(node) + # Short-circuited positions (BoolOp tails, IfExp arms) may never + # evaluate; hoisting a read out of them would fire accessors Python + # short-circuits past. A path with an unconditional occurrence + # still hoists (and its temp substitutes every occurrence). + conditional = self._conditional_eval_node_ids(node) seen: set[str] = set() results: list[tuple[str, str, str]] = [] @@ -1375,7 +3734,10 @@ def _collect_attr_reads( and isinstance(child.ctx, ast.Load) and isinstance(child.value, ast.Name) and child.value.id == "self" + and child.value.id not in bound and id(child) not in call_func_ids + and id(child) not in deferred + and id(child) not in conditional ): path_str = f"{child.value.id}.{child.attr}" if path_str not in exclude_paths and path_str not in seen: @@ -1412,12 +3774,21 @@ def _is_module_or_class_name(self, name: str) -> bool: """Return True when *name* resolves to a Python module / class / type. Looked up through ``session_data.function_globals``. Used by D1 - Name-Load instrumentation to skip module names (``cute``, ``cutlass``) + Name-Load instrumentation to skip module names and class/type references (``Boo``, ``Int32``) that should never be rewritten to ``pyir_read``. + + A local bound ONLY by plain ``import`` statements IS a module spelling + (nearest binding scope decides; other local bindings fall through to + the global view, today's answer). """ from types import ModuleType + for bound, import_only, _from_only in reversed(self._pyir_import_scope_stack): + if name in import_only: + return True + if name in bound: + break fn_globals = self.session_data.function_globals if not fn_globals: return False @@ -1430,37 +3801,54 @@ def _is_module_or_class_name(self, name: str) -> bool: return True return False + def _reading_module_name(self) -> "str | None": + """The ``__name__`` of the module whose globals resolve this function's + symbol reads; the record-only chokes root module-level bakes there.""" + fn_globals = self.session_data.function_globals + if not fn_globals: + return None + name = fn_globals.get("__name__") + return name if isinstance(name, str) else None + def _collect_other_attr_reads( - self, - node: ast.AST, - exclude_paths: set[str] | None = None, - call_func_ids: set[int] | None = None, - ) -> list[tuple[str, str, str]]: + self, node: ast.AST, exclude_paths: set[str] | None = None + ) -> list[tuple[str, str, str, bool]]: """Collect ``obj.X`` attribute reads from *node* (D1 Attribute-Load). - Returns a deduplicated list of ``(path_str, base_name, attr_name)`` - for every ``ast.Attribute(ctx=Load)`` whose ``value`` is a plain - ``ast.Name`` other than ``self``. ``self.X`` is handled by + Returns a deduplicated list of ``(path_str, base_name, attr_name, + record_only)`` for every ``ast.Attribute(ctx=Load)`` whose ``value`` + is a plain ``ast.Name`` other than ``self``. ``self.X`` is handled by :py:meth:`_collect_attr_reads`. - Only fires inside a control-flow scope -- function-scope reads - are already handled by the existing ``_mutable_ref`` / - ``_pyir_auto_load_arg`` machinery. - - *call_func_ids* lets a caller that already walked *node* pass the - call-target id set in, so we skip a redundant ``ast.walk``; it is - computed on demand when omitted. + ``record_only=False`` legs route through the staging ``pyir_read`` + hoist (in-CF reads of instance receivers, unchanged). A module/class + spelling, or a symbol-rooted read at function scope, has no staging + role but still bakes its value: those legs hoist through the + record-only ``_pyir_obs_read`` choke instead (F-SPEC observation). """ - if not self._is_inside_cf_scope(): - return [] if exclude_paths is None: exclude_paths = set() + inside_cf = self._is_inside_cf_scope() + active = self.session_data.scope_manager.get_active_symbols() + + def _in_scope(name: str) -> bool: + return any(name in scope for scope in active) - if call_func_ids is None: - call_func_ids = self._collect_call_func_ids(node) + fn_globals = self.session_data.function_globals or {} + call_func_ids = self._collect_call_func_ids(node) + # Comprehension/lambda-local names live in a private scope; must not be + # hoisted to the statement-scope prologue. + bound = _locally_bound_names(node) + # Lambda-body / genexp reads are LATE-BOUND by Python; hoisting them + # would early-bind the capture at statement time. + deferred = _deferred_execution_node_ids(node) + # Short-circuited positions (BoolOp tails, IfExp arms) may never + # evaluate; hoisting a read out of them would fire accessors Python + # short-circuits past. + conditional = self._conditional_eval_node_ids(node) seen: set[str] = set() - results: list[tuple[str, str, str]] = [] + results: list[tuple[str, str, str, bool]] = [] for child in ast.walk(node): if not ( isinstance(child, ast.Attribute) @@ -1468,21 +3856,231 @@ def _collect_other_attr_reads( and isinstance(child.value, ast.Name) ): continue - base = child.value.id - if base == "self": - continue # handled by _collect_attr_reads - if id(child) in call_func_ids: - continue - if base in ("__base_dsl__", "__module_dsl__"): + base = child.value.id + if base == "self": + continue # handled by _collect_attr_reads + if base in bound: + continue # comprehension/lambda-local; not in statement scope + if id(child) in call_func_ids: + continue + if id(child) in deferred: + continue # lambda-body / genexp read: late-bound at call time + if id(child) in conditional: + continue # short-circuited position: read stays in place + if base in ("__base_dsl__", "__module_dsl__", _PYIR_RUNTIME_ALIAS): + continue + if self._is_module_or_class_name(base): + record_only = True # module/class spelling: observe, no staging + elif inside_cf: + record_only = False # staging hoist (unchanged in-CF contract) + elif ( + not _in_scope(base) and base in fn_globals + ) or self._pyir_from_import_bound(base): + record_only = True # symbol-rooted read (global or from-import) + else: + continue # function-scope local receiver: staging owns it + path_str = f"{base}.{child.attr}" + if path_str in exclude_paths or path_str in seen: + continue + seen.add(path_str) + results.append((path_str, base, child.attr, record_only)) + return results + + @staticmethod + def _subscript_read_choke_parts(node: ast.expr) -> "tuple[str, ast.expr] | None": + """``(recorded_path, container_expr)`` when *node* is an already-emitted + ``_pyir_post_subscript_read(path, container, key)`` choke call, else None.""" + if ( + isinstance(node, ast.Call) + and isinstance(node.func, ast.Attribute) + and node.func.attr == "_pyir_post_subscript_read" + and len(node.args) == 3 + and isinstance(node.args[0], ast.Constant) + and isinstance(node.args[0].value, str) + ): + return node.args[0].value, node.args[1] + return None + + def _access_path_str(self, node: ast.expr) -> "str | None": + """Dotted path string for a PURE access-path expression (Name / Attribute / + Subscript chains and the subscript-read choke); ``None`` otherwise.""" + if isinstance(node, ast.Name): + return node.id + if isinstance(node, ast.Attribute): + base = self._access_path_str(node.value) + return None if base is None else f"{base}.{node.attr}" + if isinstance(node, ast.Subscript): + base = self._access_path_str(node.value) + if base is None: + return None + if isinstance(node.slice, ast.Constant): + return f"{base}[{node.slice.value!r}]" + key = self._access_path_str(node.slice) + return None if key is None else f"{base}[{key}]" + choke = self._subscript_read_choke_parts(node) + if choke is not None: + path, container = choke + return None if self._access_path_str(container) is None else path + return None + + def _access_path_root_name(self, node: ast.expr) -> "str | None": + """Root ``ast.Name`` id of a pure access-path expression, else None.""" + while True: + if isinstance(node, ast.Name): + return node.id + if isinstance(node, (ast.Attribute, ast.Subscript)): + node = node.value + continue + choke = self._subscript_read_choke_parts(node) + if choke is None: + return None + node = choke[1] + + @staticmethod + def _conditional_eval_node_ids(node: ast.AST) -> set[int]: + """AST node ids in positions the statement may never evaluate (BoolOp + operands after the first, IfExp arms): a hoisted pre-read of a deep + path there would evaluate what Python short-circuits past.""" + cond_ids: set[int] = set() + for child in ast.walk(node): + lazy_parts: list[ast.expr] = [] + if isinstance(child, ast.BoolOp): + lazy_parts = child.values[1:] + elif isinstance(child, ast.IfExp): + lazy_parts = [child.body, child.orelse] + for part in lazy_parts: + for sub in ast.walk(part): + cond_ids.add(id(sub)) + return cond_ids + + def _access_chain_segments( + self, node: ast.expr + ) -> "tuple[ast.expr, list[tuple[str, Any, Any]]] | None": + """Decompose a read-chain base into ``(root, hops)`` walking out from + the root. The root is an ``ast.Name`` or a call expression (hoisted + into a synthetic local at emission); each hop is ``("attr", name, + None)`` or ``("sub", key_expr, choked)`` where *choked* marks an + already-emitted subscript-read choke. Returns ``None`` for shapes + that are not access chains or whose raw subscript keys are neither + constants nor pure access paths.""" + hops: list[tuple[str, Any, Any]] = [] + while True: + if isinstance(node, ast.Name): + return node, list(reversed(hops)) + if isinstance(node, ast.Attribute): + hops.append(("attr", node.attr, None)) + node = node.value + continue + if isinstance(node, ast.Subscript): + if not ( + isinstance(node.slice, ast.Constant) + or self._access_path_str(node.slice) is not None + ): + return None + hops.append(("sub", node.slice, False)) + node = node.value + continue + if self._subscript_read_choke_parts(node) is not None: + assert isinstance(node, ast.Call) + hops.append(("sub", node.args[2], True)) + node = node.args[1] + continue + if isinstance(node, ast.Call): + return node, list(reversed(hops)) + return None + + def _collect_deep_attr_reads( + self, node: ast.AST, exclude_paths: set[str] | None = None + ) -> list[tuple[str, ast.expr, str, list[int]]]: + """Collect attribute reads whose base is itself an access chain + (``a.b.x``, ``self.d["k"].x``, ``self.helper().x``): every + ``ast.Attribute(ctx=Load)`` whose ``.value`` decomposes into hops over + a pure access path or a call-bearing root (temp-hoisted at emission). + The write side (``_slot_kwargs_for``) accepts ANY base; this hands + standalone reads the same owner/slot place fact, hop by hop. + + ``self``-rooted paths follow the ``self.X`` collector's scope rules; + other roots (including call roots) follow the ``obj.X`` collector's + CF-scope gate. Returns ``(path_str, base_node, attr_name, node_ids)`` + -- replacement is keyed on node identity because the base is an + expression, not a name. A chain deeper than the per-statement budget + refuses loudly rather than truncating. + """ + if exclude_paths is None: + exclude_paths = set() + inside_cf = self._is_inside_cf_scope() + call_func_ids = self._collect_call_func_ids(node) + # Hop receivers and keys are hoisted to statement scope; a chain + # touching a comprehension/lambda-local name cannot be hoisted. + bound = _locally_bound_names(node) + deferred = _deferred_execution_node_ids(node) + conditional = self._conditional_eval_node_ids(node) + # Only MAXIMAL chains: an Attribute serving as the base of a longer + # attribute chain is covered by that chain's read. + inner_base_ids = { + id(c.value) + for c in ast.walk(node) + if isinstance(c, ast.Attribute) and isinstance(c.value, ast.Attribute) + } + grouped: dict[str, tuple[ast.expr, str, list[int]]] = {} + for child in ast.walk(node): + if not ( + isinstance(child, ast.Attribute) + and isinstance(child.ctx, ast.Load) + and not isinstance(child.value, ast.Name) + ): + continue + if ( + id(child) in call_func_ids + or id(child) in deferred + or id(child) in conditional + or id(child) in inner_base_ids + ): continue - if self._is_module_or_class_name(base): + seg = self._access_chain_segments(child.value) + if seg is None: continue - path_str = f"{base}.{child.attr}" - if path_str in exclude_paths or path_str in seen: + root, hops = seg + if bound and any( + isinstance(sub, ast.Name) and sub.id in bound + for sub in ast.walk(child.value) + ): continue - seen.add(path_str) - results.append((path_str, base, child.attr)) - return results + if isinstance(root, ast.Name): + root_name = root.id + if root_name in ("__base_dsl__", "__module_dsl__", _PYIR_RUNTIME_ALIAS): + continue + if root_name != "self" and ( + not inside_cf or self._is_module_or_class_name(root_name) + ): + continue + base_path = self._access_path_str(child.value) + if base_path is None: + continue + path_str = f"{base_path}.{child.attr}" + if path_str in exclude_paths: + continue + group_key = path_str + else: + # Call-rooted chain: one hoisted evaluation per occurrence, + # so occurrences are never merged by spelling. + if not inside_cf: + continue + path_str = f".{child.attr}" + group_key = f"" + if len(hops) + 1 > _READ_CHAIN_HOP_BUDGET: + raise DSLUserCodeError( + DiagId.READ_DEPTH_OVERFLOW, + var=ast.unparse(child), + depth=len(hops) + 1, + budget=_READ_CHAIN_HOP_BUDGET, + ) + entry = grouped.get(group_key) + if entry is None: + grouped[group_key] = (child.value, child.attr, [id(child)]) + else: + entry[2].append(id(child)) + return [(p, b, a, ids) for p, (b, a, ids) in grouped.items()] def _is_inside_cf_scope(self) -> bool: """Return True when we are currently visiting statements that are @@ -1504,10 +4102,7 @@ def _is_inside_cf_scope(self) -> bool: return False def _collect_name_loads( - self, - node: ast.AST, - exclude_names: set[str] | None = None, - call_func_ids: set[int] | None = None, + self, node: ast.AST, exclude_names: set[str] | None = None ) -> list[str]: """Collect plain ``Name`` reads from *node* (D1 Name-Load). @@ -1531,31 +4126,29 @@ def _collect_name_loads( if exclude_names is None: exclude_names = set() - if call_func_ids is None: - call_func_ids = self._collect_call_func_ids(node) - - # Single walk: record every Attribute-base Name id (already covered by - # attribute-read instrumentation) and every Name-Load occurrence in - # source order. Iterating the collected occurrences below preserves the - # original first-occurrence result ordering without re-walking. + call_func_ids = self._collect_call_func_ids(node) + # Comprehension/lambda-local names live in a private scope; must not be + # hoisted to the statement-scope prologue. + bound = _locally_bound_names(node) + # Lambda-body / genexp reads are LATE-BOUND by Python; hoisting them + # would early-bind the capture at statement time. + deferred = _deferred_execution_node_ids(node) + # Short-circuited positions (BoolOp tails, IfExp arms) may never + # evaluate; hoisting a read out of them would evaluate what Python + # short-circuits past (an unbound name would raise eagerly). + conditional = self._conditional_eval_node_ids(node) + + # Names that appear as the .value of an Attribute (and so are + # already handled by attribute-read instrumentation) — skip them + # to keep the rewritten AST tidy. We still wrap names that ALSO + # appear elsewhere as bare reads. attr_value_only_ids: set[int] = set() - load_name_nodes: list[tuple[int, str]] = [] + bare_name_ids: set[int] = set() for child in ast.walk(node): if isinstance(child, ast.Attribute) and isinstance(child.value, ast.Name): attr_value_only_ids.add(id(child.value)) elif isinstance(child, ast.Name) and isinstance(child.ctx, ast.Load): - load_name_nodes.append((id(child), child.id)) - - # A name qualifies only if it has at least one "bare" Load occurrence — - # one that is neither an Attribute base nor a call target. Otherwise - # `boo = pyir_read('boo', boo)` would be redundant with the attribute - # instrumentation for `boo.val`. Precomputing this set replaces the - # per-candidate ``ast.walk`` rescan the loop used to do. - names_with_bare_load = { - name - for node_id, name in load_name_nodes - if node_id not in attr_value_only_ids and node_id not in call_func_ids - } + bare_name_ids.add(id(child)) active = self.session_data.scope_manager.get_active_symbols() @@ -1564,12 +4157,30 @@ def _in_scope(name: str) -> bool: seen: set[str] = set() results: list[str] = [] - for node_id, name in load_name_nodes: - if node_id in call_func_ids: + for child in ast.walk(node): + if not (isinstance(child, ast.Name) and isinstance(child.ctx, ast.Load)): + continue + if id(child) in call_func_ids: continue + if id(child) in deferred: + continue # lambda-body / genexp read: late-bound at call time + if id(child) in conditional: + continue # short-circuited position: read stays in place + if ( + id(child) in attr_value_only_ids + and id(child) not in bare_name_ids - attr_value_only_ids + ): + # The only occurrences are as Attribute bases — the + # attribute-read instrumentation already covers reads + # through obj.attr. Wrapping the bare obj would be + # redundant. + pass # but we still allow if it appears as bare elsewhere + name = child.id if name in exclude_names or name in seen: continue - if name in ("__base_dsl__", "__module_dsl__"): + if name in bound: + continue # comprehension/lambda-local; not in statement scope + if name in ("__base_dsl__", "__module_dsl__", _PYIR_RUNTIME_ALIAS): continue if self.session_data.scope_manager.is_skip_reference_taking(name): continue @@ -1577,8 +4188,93 @@ def _in_scope(name: str) -> bool: continue if self._is_module_or_class_name(name): continue - if name not in names_with_bare_load: + # Skip names that ONLY appear as Attribute base. Without this + # we emit `boo = pyir_read('boo', boo)` for every `boo.val`, + # which is correct but noisy. + other_uses = [ + c + for c in ast.walk(node) + if isinstance(c, ast.Name) + and isinstance(c.ctx, ast.Load) + and c.id == name + and id(c) not in attr_value_only_ids + and id(c) not in call_func_ids + and id(c) not in deferred + and id(c) not in conditional + ] + if not other_uses: + continue + seen.add(name) + results.append(name) + return results + + def _collect_global_name_reads( + self, node: ast.AST, exclude_names: set[str] | None = None + ) -> list[str]: + """Collect bare GLOBAL ``Name(ctx=Load)`` reads from *node*: names not + bound in any active local scope that resolve through the function's + globals to a non-module/class/function value. These reads have no + staging choke at any scope, so they hoist through the record-only + ``_pyir_obs_global_read`` choke (F-SPEC observation).""" + if exclude_names is None: + exclude_names = set() + fn_globals = self.session_data.function_globals + if not fn_globals: + return [] + active = self.session_data.scope_manager.get_active_symbols() + + def _in_scope(name: str) -> bool: + return any(name in scope for scope in active) + + call_func_ids = self._collect_call_func_ids(node) + bound = _locally_bound_names(node) + deferred = _deferred_execution_node_ids(node) + # Short-circuited positions (BoolOp tails, IfExp arms) may never + # evaluate; hoisting a read out of them would evaluate what Python + # short-circuits past (a missing global would raise eagerly). + conditional = self._conditional_eval_node_ids(node) + # Names serving only as Attribute bases are observed by the attr choke. + attr_base_ids = { + id(child.value) + for child in ast.walk(node) + if isinstance(child, ast.Attribute) and isinstance(child.value, ast.Name) + } + + seen: set[str] = set() + results: list[str] = [] + for child in ast.walk(node): + if not (isinstance(child, ast.Name) and isinstance(child.ctx, ast.Load)): + continue + name = child.id + if name in exclude_names or name in seen: + continue + if id(child) in call_func_ids or id(child) in deferred: + continue + if id(child) in conditional: + continue # short-circuited position: read stays in place + if id(child) in attr_base_ids: + continue + if name in bound or name.startswith("_pyir_"): + continue + if name in ("__base_dsl__", "__module_dsl__", _PYIR_RUNTIME_ALIAS): + continue + if self.session_data.scope_manager.is_skip_reference_taking(name): continue + if _in_scope(name): + continue # a local binding, never the module global + if name not in fn_globals: + continue # builtin / closure free var: no module root + obj = fn_globals.get(name) + if isinstance( + obj, + ( + types.ModuleType, + type, + types.FunctionType, + types.BuiltinFunctionType, + ), + ): + continue # symbol references, never scalar bakes seen.add(name) results.append(name) return results @@ -1618,10 +4314,60 @@ def _build_first_def_pyir_assigns( slot_kwargs = self._slot_kwargs_for(target) target_load = self._target_as_load(target) + written_kwargs = [_deepcopy_ast_root(kw) for kw in slot_kwargs] + # Declare a CONSTRUCTED container RHS (literal/comprehension/concat): + # fresh, so element slots may be replaced in place. Alias RHS: no fact. + if isinstance(target, ast.Name) and isinstance( + node.value, + ( + ast.List, + ast.Dict, + ast.Set, + ast.ListComp, + ast.SetComp, + ast.DictComp, + ast.BinOp, + ), + ): + written_kwargs = written_kwargs + [ + ast.keyword(arg="fresh_binding", value=ast.Constant(value=True)) + ] + # Declare constant-keyed NESTED constructions so the re-init recurses + # exactly into re-constructed sub-containers (never an alias entry). + _fresh_paths = self._nested_fresh_container_paths(node.value) + if _fresh_paths: + written_kwargs = written_kwargs + [ + ast.keyword( + arg="fresh_paths", + value=ast.Tuple( + elts=[ + ast.Tuple( + elts=[ast.Constant(value=_c) for _c in _p], + ctx=ast.Load(), + ) + for _p in _fresh_paths + ], + ctx=ast.Load(), + ), + ) + ] + # Declare a direct-call RHS's callee so the runtime can witness a + # plain-allocating construction (a chain via the new binding never can). + if ( + isinstance(target, ast.Name) + and isinstance(node.value, ast.Call) + and self._callee_root_name(node.value.func) not in (None, target.id) + ): + written_kwargs = written_kwargs + [ + ast.keyword( + arg="rhs_ctor", + value=self._target_as_load(node.value.func), + ) + ] + pyir_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "pyir_assign", - submodule_name="pyir_runtime", lineno=lineno, col_offset=node.col_offset, ), @@ -1632,7 +4378,7 @@ def _build_first_def_pyir_assigns( ast.Constant(value=self.session_data.file_name), ast.Constant(value=lineno), ], - keywords=[deepcopy(kw) for kw in slot_kwargs], + keywords=written_kwargs, ) reassign = ast.Assign( targets=[_deepcopy_ast_root(target)], @@ -1641,6 +4387,51 @@ def _build_first_def_pyir_assigns( stmts.append(ast.copy_location(ast.fix_missing_locations(reassign), node)) return stmts + def _build_generation_probe_stmts( + self, node: ast.stmt, exclude_paths: set[str] | None = None + ) -> list[ast.stmt]: + """One ``pyir_generation_probe("base.attr")`` per Name-rooted attribute load + of a function-scope statement; empty inside a CF scope.""" + if self._is_inside_cf_scope(): + return [] + call_func_ids = self._collect_call_func_ids(node) + bound = _locally_bound_names(node) + seen: set[str] = set() + stmts: list[ast.stmt] = [] + for child in ast.walk(node): + if not ( + isinstance(child, ast.Attribute) + and isinstance(child.ctx, ast.Load) + and isinstance(child.value, ast.Name) + ): + continue + base = child.value.id + if base in bound or id(child) in call_func_ids: + continue + if base in ("__base_dsl__", "__module_dsl__", _PYIR_RUNTIME_ALIAS): + continue + if self._is_module_or_class_name(base): + continue + path_str = f"{base}.{child.attr}" + if (exclude_paths and path_str in exclude_paths) or path_str in seen: + continue + seen.add(path_str) + lineno = node.lineno + col_offset = getattr(node, "col_offset", 0) + probe = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "pyir_generation_probe", + lineno=lineno, + col_offset=col_offset, + ), + args=[ast.Constant(value=path_str)], + keywords=[], + ) + ) + stmts.append(ast.copy_location(ast.fix_missing_locations(probe), node)) + return stmts + def _insert_pyir_attr_reads( self, node: ast.stmt, exclude_paths: set[str] | None = None ) -> ast.stmt | list[ast.stmt]: @@ -1667,31 +4458,44 @@ def _insert_pyir_attr_reads( Returns ``[read_stmt, ..., modified_node]`` if any pattern fired, else *node* unchanged. All three patterns run on the same node so - a single ``cute.printf(boo.val, bar)`` statement instruments both + a single downstream-DSL print call ``print_api(boo.val, bar)`` statement instruments both ``boo.val`` and ``bar`` in one call. """ - # All three collectors need the call-target id set for *node*; compute - # it once here instead of three identical ``ast.walk`` passes. - call_func_ids = self._collect_call_func_ids(node) - self_reads = self._collect_attr_reads(node, exclude_paths, call_func_ids) - other_attr_reads = self._collect_other_attr_reads( - node, exclude_paths, call_func_ids - ) - name_loads = self._collect_name_loads(node, exclude_paths, call_func_ids) + exclude_paths_set = exclude_paths or set() + self_reads = self._collect_attr_reads(node, exclude_paths) + other_attr_reads = self._collect_other_attr_reads(node, exclude_paths) + deep_attr_reads = self._collect_deep_attr_reads(node, exclude_paths) + name_loads = self._collect_name_loads(node, exclude_paths) + global_name_reads = self._collect_global_name_reads(node, exclude_paths) # Don't emit a Name-Load writeback for a name that already serves # as the owner of a non-self Attribute writeback we just emitted: # `boo = pyir_read('boo', boo)` is redundant with # `boo.val = pyir_read('boo.val', boo.val, owner=boo, ...)`. - attr_owner_names = {base for _, base, _ in other_attr_reads} + attr_owner_names = {base for _, base, _, _ in other_attr_reads} name_loads = [n for n in name_loads if n not in attr_owner_names] - if not self_reads and not other_attr_reads and not name_loads: + # Function-scope superseded-generation probes (empty inside CF scope, + # where the collectors above carry the access-root name instead). + probe_stmts = self._build_generation_probe_stmts(node, exclude_paths) + + if ( + not self_reads + and not other_attr_reads + and not deep_attr_reads + and not name_loads + and not global_name_reads + ): + if probe_stmts: + return [*probe_stmts, node] return node stmts: list[ast.stmt] = [] # Map from (base_name, attr_name) -> local var name for replacement replacements: dict[tuple[str, str], str] = {} + # Path -> temp of every read emitted for this statement; deep chains + # reuse these as hop receivers instead of re-reading a prefix. + hop_temps: dict[str, str] = {} # ----- (1) self.X reads using local-var pattern (existing) ----- for path_str, base_name, attr_name in self_reads: @@ -1701,6 +4505,7 @@ def _insert_pyir_attr_reads( local_var = f"_pyir_attr_{self.session_data.counter}" self.session_data.counter += 1 replacements[(base_name, attr_name)] = local_var + hop_temps[path_str] = local_var attr_load = ast.Attribute( value=ast.Name(id=base_name, ctx=ast.Load()), @@ -1708,9 +4513,8 @@ def _insert_pyir_attr_reads( ctx=ast.Load(), ) pyir_read_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "pyir_read", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), @@ -1744,42 +4548,61 @@ def _insert_pyir_attr_reads( # mirrors the self.X behaviour, keeping ``_mutable_ref`` from # leaking through object boundaries when the load is fed into # constructors or returns. - for path_str, base_name, attr_name in other_attr_reads: + for path_str, base_name, attr_name, record_only in other_attr_reads: lineno = node.lineno col_offset = getattr(node, "col_offset", 0) local_var = f"_pyir_attr_{self.session_data.counter}" self.session_data.counter += 1 replacements[(base_name, attr_name)] = local_var - - attr_load = ast.Attribute( - value=ast.Name(id=base_name, ctx=ast.Load()), - attr=attr_name, - ctx=ast.Load(), - ) - pyir_read_call = ast.Call( - func=_create_module_attribute( - "pyir_read", - submodule_name="pyir_runtime", - lineno=lineno, - col_offset=col_offset, - ), - args=[ - ast.Constant(value=path_str), - attr_load, - ], - keywords=[ - ast.keyword(arg="attach_ref", value=ast.Constant(value=False)), - ast.keyword( - arg="owner", - value=ast.Name(id=base_name, ctx=ast.Load()), + hop_temps[path_str] = local_var + + if record_only: + # Symbol-rooted read with no staging role (module/class + # spelling or function-scope global object): the record-only + # observation choke evaluates the same getattr in place. + pyir_read_call = ast.Call( + func=_create_runtime_attribute( + "_pyir_obs_read", + lineno=lineno, + col_offset=col_offset, ), - ast.keyword( - arg="slot_name", - value=ast.Constant(value=attr_name), + args=[ + ast.Constant(value=path_str), + ast.Name(id=base_name, ctx=ast.Load()), + ast.Constant(value=attr_name), + ast.Constant(value=self._reading_module_name()), + ], + keywords=[], + ) + else: + attr_load = ast.Attribute( + value=ast.Name(id=base_name, ctx=ast.Load()), + attr=attr_name, + ctx=ast.Load(), + ) + pyir_read_call = ast.Call( + func=_create_runtime_attribute( + "pyir_read", + lineno=lineno, + col_offset=col_offset, ), - ], - ) + args=[ + ast.Constant(value=path_str), + attr_load, + ], + keywords=[ + ast.keyword(arg="attach_ref", value=ast.Constant(value=False)), + ast.keyword( + arg="owner", + value=ast.Name(id=base_name, ctx=ast.Load()), + ), + ast.keyword( + arg="slot_name", + value=ast.Constant(value=attr_name), + ), + ], + ) read_stmt = ast.Assign( targets=[ast.Name(id=local_var, ctx=ast.Store())], value=pyir_read_call, @@ -1792,7 +4615,7 @@ def _insert_pyir_attr_reads( # reasons: # * Write-back inside a closure body shadows the outer name, # forcing it into the loop's ``write_args``. Read-only loop - # variables like ``tar`` in mixing.py would then be threaded + # variables like ``tar`` in mixing.py would then be carried # through SCF and emerge as dynamic ``Boolean`` SSA, breaking # ``const_expr(tar)``. # * Wrapping a closure-captured function-arg (``n`` in @@ -1807,6 +4630,10 @@ def _insert_pyir_attr_reads( # path (for Python primitives) -- DSL values are passed through # untouched, and the existing ``_pyir_auto_load_arg`` boundary # still emits ``pyir.load`` for values that DO carry a ref. + # + # Emitted BEFORE the deep-chain hoists: a hoisted call base + # references the name temps, and a local name binding cannot be + # changed by evaluating the call, so the pre-read is safe. name_replacements: dict[str, str] = {} for name in name_loads: lineno = node.lineno @@ -1817,9 +4644,8 @@ def _insert_pyir_attr_reads( name_replacements[name] = local_var pyir_read_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "pyir_read", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), @@ -1837,17 +4663,184 @@ def _insert_pyir_attr_reads( ) stmts.append(ast.copy_location(ast.fix_missing_locations(read_stmt), node)) + # ----- (3b) bare GLOBAL Name reads: record-only observation ----- + # No staging choke exists for a module-level binding at any scope; + # the baked value still records under its module root (F-SPEC). + for name in global_name_reads: + if name in name_replacements: + continue + lineno = node.lineno + col_offset = getattr(node, "col_offset", 0) + + local_var = f"_pyir_name_{self.session_data.counter}" + self.session_data.counter += 1 + name_replacements[name] = local_var + + obs_call = ast.Call( + func=_create_runtime_attribute( + "_pyir_obs_global_read", + lineno=lineno, + col_offset=col_offset, + ), + args=[ + ast.Constant(value=name), + ast.Name(id=name, ctx=ast.Load()), + ast.Constant(value=self._reading_module_name()), + ], + keywords=[], + ) + read_stmt = ast.Assign( + targets=[ast.Name(id=local_var, ctx=ast.Store())], + value=obs_call, + ) + stmts.append(ast.copy_location(ast.fix_missing_locations(read_stmt), node)) + + # ----- (2b) deep-path reads: hop-wise three-address lowering ----- + # Each hop is choked on its runtime-resolved receiver (owner = the + # previous hop's temp), so every leg of a.b["k"].x lands on the place + # the write side names; a call-bearing base hoists the call into a + # temp first (single evaluation) and later hops chain off the temp. + deep_replacements: dict[int, str] = {} + call_hoist_stmts: list[ast.stmt] = [] + if deep_attr_reads: + # Chains nested inside another chain's base (call args, keys) + # emit first so the enclosing hoist references their temps. + node_depths = _ast_node_depths(node) + deep_attr_reads = sorted( + deep_attr_reads, + key=lambda rec: -max(node_depths.get(nid, 0) for nid in rec[3]), + ) + for _, base_node, attr_name, node_ids in deep_attr_reads: + lineno = node.lineno + col_offset = getattr(node, "col_offset", 0) + seg = self._access_chain_segments(base_node) + if seg is None: + continue + root, hops = seg + hops = [*hops, ("attr", attr_name, None)] + if isinstance(root, ast.Name): + recv = root.id + cur_path = root.id + else: + recv = f"_pyir_call_{self.session_data.counter}" + self.session_data.counter += 1 + # The call node moves out of the statement verbatim (its + # occurrence is replaced wholesale by the final hop temp). + hoist = ast.Assign( + targets=[ast.Name(id=recv, ctx=ast.Store())], + value=root, + ) + hoist = ast.copy_location(ast.fix_missing_locations(hoist), node) + stmts.append(hoist) + call_hoist_stmts.append(hoist) + cur_path = recv + for kind, payload, choked in hops: + if kind == "attr": + hop_path = f"{cur_path}.{payload}" + elif isinstance(payload, ast.Constant): + hop_path = f"{cur_path}[{payload.value!r}]" + else: + hop_path = f"{cur_path}[{self._access_path_str(payload)}]" + existing = hop_temps.get(hop_path) + if existing is not None: + recv, cur_path = existing, hop_path + continue + local_var = f"_pyir_attr_{self.session_data.counter}" + self.session_data.counter += 1 + recv_load = ast.Name(id=recv, ctx=ast.Load()) + value: ast.expr + if kind == "attr": + attr_load = ast.Attribute( + value=recv_load, attr=payload, ctx=ast.Load() + ) + if hop_path in exclude_paths_set: + # Excluded place: plain three-address leg, no choke. + value = attr_load + else: + value = ast.Call( + func=_create_runtime_attribute( + "pyir_read", + lineno=lineno, + col_offset=col_offset, + ), + args=[ast.Constant(value=hop_path), attr_load], + keywords=[ + ast.keyword( + arg="attach_ref", + value=ast.Constant(value=False), + ), + ast.keyword( + arg="owner", + value=ast.Name(id=recv, ctx=ast.Load()), + ), + ast.keyword( + arg="slot_name", + value=ast.Constant(value=payload), + ), + ], + ) + else: + key_node = self._target_as_load(payload) + if choked and hop_path not in exclude_paths_set: + value = ast.Call( + func=_create_runtime_attribute( + "_pyir_post_subscript_read", + lineno=lineno, + col_offset=col_offset, + ), + args=[ + ast.Constant(value=hop_path), + recv_load, + key_node, + ], + keywords=[], + ) + else: + value = ast.Subscript( + value=recv_load, slice=key_node, ctx=ast.Load() + ) + read_stmt = ast.Assign( + targets=[ast.Name(id=local_var, ctx=ast.Store())], + value=value, + ) + stmts.append( + ast.copy_location(ast.fix_missing_locations(read_stmt), node) + ) + hop_temps[hop_path] = local_var + recv, cur_path = local_var, hop_path + for nid in node_ids: + deep_replacements[nid] = recv + # Replace self.X / non-self Attribute references with their # local vars, then bare Name references with their local vars. - replaced = self._replace_attr_with_local(node, replacements) + # Deferred regions (lambda bodies / genexp lazy parts) keep their + # original reads so the capture stays late-bound (Python semantics). + deferred = _deferred_execution_node_ids(node) + replaced = self._replace_attr_with_local( + node, replacements, deferred, deep_replacements + ) if name_replacements: - replaced = self._replace_name_with_local(replaced, name_replacements) + replaced = self._replace_name_with_local( + replaced, name_replacements, deferred + ) + # A hoisted call base moved out of the statement before replacement + # ran; its argument reads route through the same temps (node ids are + # preserved by the move, so id-keyed replacement still applies). + for hoist_stmt in call_hoist_stmts: + self._replace_attr_with_local( + hoist_stmt, replacements, deferred, deep_replacements + ) + if name_replacements: + self._replace_name_with_local(hoist_stmt, name_replacements, deferred) assert isinstance(replaced, ast.stmt) stmts.append(replaced) return stmts def _replace_name_with_local( - self, node: ast.AST, replacements: dict[str, str] + self, + node: ast.AST, + replacements: dict[str, str], + deferred: set[int] | None = None, ) -> ast.AST: """Replace bare ``Name(ctx=Load)`` references in *node* with ``Name(local_var, ctx=Load)``. Used by D1 Name-Load @@ -1856,13 +4849,16 @@ def _replace_name_with_local( Does NOT descend into ``ast.Attribute``'s ``.value`` slot -- ``obj.X`` writes/reads on ``obj`` are already handled by the attribute-read pass. Does NOT replace Store-context Names. + Nodes in *deferred* (lambda bodies / genexp lazy parts) are left + untouched so their evaluation stays late-bound. """ if not replacements: return node + deferred_ids = deferred or set() class NameReplacer(ast.NodeTransformer): def visit_Name(self, child: ast.Name) -> ast.AST: - if isinstance(child.ctx, ast.Load): + if isinstance(child.ctx, ast.Load) and id(child) not in deferred_ids: new_id = replacements.get(child.id) if new_id is not None: return ast.copy_location( @@ -1880,22 +4876,38 @@ def visit_Attribute(self, child: ast.Attribute) -> ast.AST: return NameReplacer().visit(node) def _replace_attr_with_local( - self, node: ast.AST, replacements: dict[tuple[str, str], str] + self, + node: ast.AST, + replacements: dict[tuple[str, str], str], + deferred: set[int] | None = None, + deep_replacements: dict[int, str] | None = None, ) -> ast.AST: """Replace ``self.X`` attribute reads in *node* with local variable names from *replacements*. *replacements* maps ``(base_name, attr_name)`` → ``local_var_name``. Only replaces ``ast.Attribute(ctx=Load)`` with matching base Name. + *deep_replacements* maps the ``id()`` of a collected deep-path + Attribute node → local var (deep bases are expressions, not names). + Nodes in *deferred* (lambda bodies / genexp lazy parts) are left + untouched so their evaluation stays late-bound. """ - if not replacements: + deep_ids = deep_replacements or {} + if not replacements and not deep_ids: return node + deferred_ids = deferred or set() class AttrReplacer(ast.NodeTransformer): def visit_Attribute(self, child: ast.Attribute) -> ast.AST: + if id(child) in deep_ids and id(child) not in deferred_ids: + return ast.copy_location( + ast.Name(id=deep_ids[id(child)], ctx=ast.Load()), child + ) self.generic_visit(child) - if isinstance(child.ctx, ast.Load) and isinstance( - child.value, ast.Name + if ( + isinstance(child.ctx, ast.Load) + and isinstance(child.value, ast.Name) + and id(child) not in deferred_ids ): key = (child.value.id, child.attr) local_var = replacements.get(key) @@ -1925,19 +4937,94 @@ def _insert_pyir_assign( (e.g. ``self._idx = idx`` in ``__init__``) are skipped safely. """ stmts: list[ast.stmt] = [] - for target in targets: + + # A chained assignment (``a = b = expr``) evaluates its RHS ONCE; hoist + # it to a shared temp so each target assigns from the temp exactly once. + chain_name: "str | None" = None + if len(node.targets) > 1 and len(targets) > 1: + chain_name = f"_pyir_chain_{self.session_data.counter}" + self.session_data.counter += 1 + rhs_hoist = ast.Assign( + targets=[ast.Name(id=chain_name, ctx=ast.Store())], + value=node.value, + ) + stmts.append(ast.copy_location(ast.fix_missing_locations(rhs_hoist), node)) + + node_reinserted = False + + def _orig_stmt_for(tgt: ast.expr) -> ast.stmt: + """The 'original assignment' statement for *tgt*: the whole node + in the single-eval case, else ``tgt = ``. + + The returned statements take the tree position *node* vacates, + so the first single-eval insertion reuses *node* itself; only a + second insertion (the hasattr guard and first-def fallback + branches of an Attribute target) copies, keeping every tree + position a distinct node.""" + nonlocal node_reinserted + if chain_name is None: + if node_reinserted: + return _deepcopy_ast_root(node) + node_reinserted = True + return node + single = ast.Assign( + targets=[_deepcopy_ast_root(tgt)], + value=ast.Name(id=chain_name, ctx=ast.Load()), + ) + return ast.copy_location(ast.fix_missing_locations(single), node) + + iter_targets = list(node.targets) if chain_name is not None else targets + for target in iter_targets: + if chain_name is not None and not any(t is target for t in targets): + # Uninstrumented sibling of a chained assign (e.g. a first-def): + # assign it from the shared temp exactly once. + stmts.append(_orig_stmt_for(target)) + continue path_str = self._target_to_path_str(target) lineno = node.lineno slot_kwargs = self._slot_kwargs_for(target) + # An Attribute target routes both instrumentation calls through + # the attr-assign wrappers, which return ``_PYIR_SKIP`` when the + # owner keeps handle-shaped state behind a custom ``__setattr__`` + # (a downstream-DSL struct field reads back as its scalar-slot + # pointer handle -- writing that handle back through the real + # ``__setattr__`` raises). The generated shape mirrors the + # subscript pre-hook protocol: capture into a temp, store only + # when the temp is not ``_PYIR_SKIP``. Every other target keeps + # the plain single-statement shape. + is_attr_target = isinstance(target, ast.Attribute) + + def _skip_guarded_store(tmp_name: str) -> ast.stmt: + """``if is not _PYIR_SKIP: target = ``.""" + return ast.If( + test=ast.Compare( + left=ast.Name(id=tmp_name, ctx=ast.Load()), + ops=[ast.IsNot()], + comparators=[ + _create_runtime_attribute( + "_PYIR_SKIP", + lineno=lineno, + col_offset=node.col_offset, + ) + ], + ), + body=[ + ast.Assign( + targets=[_deepcopy_ast_root(target)], + value=ast.Name(id=tmp_name, ctx=ast.Load()), + ) + ], + orelse=[], + ) + # target = pyir_read("target", target, owner=..., slot_name=...) # For dict-style subscripts (d['key']), use d.get('key') to # avoid KeyError on first-time key insertion. target_load_for_read = self._subscript_safe_read(target) pyir_read_call = ast.Call( - func=_create_module_attribute( - "pyir_read", - submodule_name="pyir_runtime", + func=_create_runtime_attribute( + "_pyir_pre_attr_assign" if is_attr_target else "pyir_read", lineno=lineno, col_offset=node.col_offset, ), @@ -1945,12 +5032,25 @@ def _insert_pyir_assign( ast.Constant(value=path_str), target_load_for_read, ], - keywords=[deepcopy(kw) for kw in slot_kwargs], - ) - read_stmt = ast.Assign( - targets=[deepcopy(target)], - value=pyir_read_call, + keywords=[_deepcopy_ast_root(kw) for kw in slot_kwargs], ) + if is_attr_target: + read_tmp_name = f"_pyir_read_{self.session_data.counter}" + self.session_data.counter += 1 + read_stmts: list[ast.stmt] = [ + ast.Assign( + targets=[ast.Name(id=read_tmp_name, ctx=ast.Store())], + value=pyir_read_call, + ), + _skip_guarded_store(read_tmp_name), + ] + else: + read_stmts = [ + ast.Assign( + targets=[_deepcopy_ast_root(target)], + value=pyir_read_call, + ) + ] # _old = target old_name = f"_pyir_old_{self.session_data.counter}" @@ -1965,9 +5065,8 @@ def _insert_pyir_assign( # owner=..., slot_name=...) target_load2 = self._target_as_load(target) pyir_call = ast.Call( - func=_create_module_attribute( - "pyir_assign", - submodule_name="pyir_runtime", + func=_create_runtime_attribute( + "_pyir_post_attr_assign" if is_attr_target else "pyir_assign", lineno=lineno, col_offset=node.col_offset, ), @@ -1978,83 +5077,294 @@ def _insert_pyir_assign( ast.Constant(value=self.session_data.file_name), ast.Constant(value=lineno), ], - keywords=[deepcopy(kw) for kw in slot_kwargs], - ) - reassign = ast.Assign( - targets=[deepcopy(target)], - value=pyir_call, + keywords=[_deepcopy_ast_root(kw) for kw in slot_kwargs], ) + if is_attr_target: + assign_tmp_name = f"_pyir_assign_{self.session_data.counter}" + self.session_data.counter += 1 + reassign_stmts: list[ast.stmt] = [ + ast.Assign( + targets=[ast.Name(id=assign_tmp_name, ctx=ast.Store())], + value=pyir_call, + ), + _skip_guarded_store(assign_tmp_name), + ] + else: + reassign_stmts = [ + ast.Assign( + targets=[_deepcopy_ast_root(target)], + value=pyir_call, + ) + ] # Attribute targets might be first definitions (e.g. # self._idx = idx in __init__). Guard with hasattr so # the pyir_read of the old value doesn't crash when the # attribute doesn't exist yet. Applies in both callee - # rewrite AND top-level @cute.jit paths. - if isinstance(target, ast.Attribute) and isinstance(target.value, ast.Name): - # if hasattr(obj, attr): - # pyir_read + capture_old + original + pyir_assign - # else: - # original + # rewrite AND top-level jit paths. + if isinstance(target, ast.Attribute): + # Guard the pre-read on hasattr(obj, attr); obj may itself be a + # nested load (self.inner.newattr), so load the value chain. hasattr_test = ast.Call( func=ast.Name(id="hasattr", ctx=ast.Load()), args=[ - ast.Name(id=target.value.id, ctx=ast.Load()), + self._target_as_load(target.value), ast.Constant(value=target.attr), ], keywords=[], ) + # Tracer-internal storage probe, not a user reflection read: + # keep it bare through the call-boundary pass. + hasattr_test._pyir_synth = True # type: ignore[attr-defined] guarded_body: list[ast.stmt] = [ - ast.copy_location(ast.fix_missing_locations(read_stmt), node), + *( + ast.copy_location(ast.fix_missing_locations(s), node) + for s in read_stmts + ), ast.copy_location(ast.fix_missing_locations(capture_old), node), - deepcopy(node), - ast.copy_location(ast.fix_missing_locations(reassign), node), + _orig_stmt_for(target), + *( + ast.copy_location(ast.fix_missing_locations(s), node) + for s in reassign_stmts + ), + ] + # First-def branch: the plain assignment plus a trace-only note + # recording the binding-position fact (read-before-set). + note_call = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "pyir_note_attr_first_def", + lineno=lineno, + col_offset=node.col_offset, + ), + args=[ + ast.Constant(value=path_str), + self._target_as_load(target.value), + ast.Constant(value=target.attr), + self._target_as_load(target), + ], + keywords=[], + ) + ) + fallback_body: list[ast.stmt] = [ + _orig_stmt_for(target), + ast.copy_location(ast.fix_missing_locations(note_call), node), ] - fallback_body: list[ast.stmt] = [deepcopy(node)] - if_guard = ast.If( + if_guard: ast.stmt = ast.If( test=hasattr_test, body=guarded_body, orelse=fallback_body, ) + # A property target dispatches its SETTER through the call + # boundary (class code, no ast.Call): the native store pattern + # would re-invoke the setter per bracketing statement and + # promote a placeless (computed) leg. The branch fact is + # value-free, so the RHS keeps its native evaluation position + # (after the pre-store refresh) in whichever branch runs. A + # multi-target assign without a chain temp re-executes the + # whole node, so it keeps the storage-only shape. + if chain_name is not None or len(node.targets) == 1: + is_prop_test = ast.Call( + func=_create_runtime_attribute( + "_pyir_is_property_store", + lineno=lineno, + col_offset=node.col_offset, + ), + args=[ + self._target_as_load(target.value), + ast.Constant(value=target.attr), + ], + keywords=[], + ) + is_prop_test._pyir_synth = True # type: ignore[attr-defined] + prop_store = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_property_store", + lineno=lineno, + col_offset=node.col_offset, + ), + args=[ + self._target_as_load(target.value), + ast.Constant(value=target.attr), + ( + ast.Name(id=chain_name, ctx=ast.Load()) + if chain_name is not None + else _deepcopy_ast_root(node.value) + ), + ], + keywords=[], + ) + ) + if_guard = ast.If( + test=is_prop_test, + body=[ + ast.copy_location( + ast.fix_missing_locations(prop_store), node + ) + ], + orelse=[if_guard], + ) stmts.append( ast.copy_location(ast.fix_missing_locations(if_guard), node) ) else: - stmts.append( - ast.copy_location(ast.fix_missing_locations(read_stmt), node) + # A statically in-scope Name can be UNBOUND at run time (prior + # ``del`` / failed bind): seed ``_pyir_old = None`` on unbound read. + set_old_none = ast.Assign( + targets=[ast.Name(id=old_name, ctx=ast.Store())], + value=ast.Constant(value=None), ) - stmts.append( - ast.copy_location(ast.fix_missing_locations(capture_old), node) + guarded_read = ast.Try( + body=[ + *( + ast.copy_location(ast.fix_missing_locations(s), node) + for s in read_stmts + ), + ast.copy_location(ast.fix_missing_locations(capture_old), node), + ], + handlers=[ + ast.ExceptHandler( + type=ast.Tuple( + elts=[ + ast.Name(id="NameError", ctx=ast.Load()), + ast.Name(id="UnboundLocalError", ctx=ast.Load()), + ], + ctx=ast.Load(), + ), + name=None, + body=[ + ast.copy_location( + ast.fix_missing_locations(set_old_none), node + ) + ], + ) + ], + orelse=[], + finalbody=[], ) - stmts.append(node) # original assignment stmts.append( - ast.copy_location(ast.fix_missing_locations(reassign), node) + ast.copy_location(ast.fix_missing_locations(guarded_read), node) + ) + stmts.append(_orig_stmt_for(target)) # original assignment + stmts.extend( + ast.copy_location(ast.fix_missing_locations(s), node) + for s in reassign_stmts ) if len(stmts) == 0: return node return stmts - def _insert_pyir_augassign(self, node: ast.AugAssign) -> list[ast.stmt]: + # ``ast.AugAssign`` operator -> the in-place routing key consumed by + # ``_pyir_inplace_binop`` (total over Python's binary operators). + _AUGASSIGN_OP_KEYS: "dict[type, str]" = { + ast.Add: "add", + ast.Sub: "sub", + ast.Mult: "mul", + ast.MatMult: "matmul", + ast.Div: "truediv", + ast.FloorDiv: "floordiv", + ast.Mod: "mod", + ast.Pow: "pow", + ast.LShift: "lshift", + ast.RShift: "rshift", + ast.BitOr: "or", + ast.BitXor: "xor", + ast.BitAnd: "and", + } + + def _inplace_binop_stmt( + self, node: ast.AugAssign, base_name: "str | None" = None + ) -> ast.stmt: + """``target op= value`` -> ``target = _pyir_inplace_binop("op", target, + value)``: value-identical (CPython's in-place protocol), with the dunder + routed through the call boundary (operator dispatch has no ast.Call). + *base_name* substitutes a pre-walrus anchor temp for the target load + (Python loads the base before a same-statement walrus store).""" + base_load: ast.expr = ( + ast.Name(id=base_name, ctx=ast.Load()) + if base_name is not None + else self._target_as_load(node.target) + ) + call = ast.Call( + func=_create_runtime_attribute( + "_pyir_inplace_binop", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[ + ast.Constant(value=self._AUGASSIGN_OP_KEYS[type(node.op)]), + base_load, + _deepcopy_ast_root(node.value), + ], + keywords=[], + ) + stmt = ast.Assign(targets=[_deepcopy_ast_root(node.target)], value=call) + return ast.copy_location(ast.fix_missing_locations(stmt), node) + + def _insert_pyir_augassign( + self, node: ast.AugAssign, base_anchor: "str | None" = None + ) -> list[ast.stmt]: """Insert pyir_read + pyir_assign around an augmented assignment. Transforms: target += expr To: target = pyir_read(name, target) # load from ref _old = target - target += expr # uses loaded value + target = _pyir_inplace_binop(op, target, expr) target = pyir_assign(name, _old, target, file, line) + + *base_anchor* (a Name target walrus-rebound in its own RHS) replaces + the in-place op's target load with the statement-top anchor temp; the + pyir_read refresh and old-capture keep the post-walrus slot state for + the store. + + A property target instead branch-selects on the same value-free + target fact as plain assigns: one getter read feeds the in-place + op and the bound setter runs once at the call boundary (the + storage pattern would re-invoke the accessors per bracketing + statement and promote a placeless leg). """ target = node.target path_str = self._target_to_path_str(target) lineno = node.lineno slot_kwargs = self._slot_kwargs_for(target) + # Attribute targets route through the attr-assign wrappers and the + # ``_PYIR_SKIP``-guarded store, exactly like ``_insert_pyir_assign`` + # (see the comment there): a handle-shaped owner (a downstream-DSL + # struct whose fields read back as pointer handles) must not receive + # its own read-back through the real __setattr__. + is_attr_target = isinstance(target, ast.Attribute) + + def _skip_guarded_store(tmp_name: str) -> ast.stmt: + """``if is not _PYIR_SKIP: target = ``.""" + return ast.If( + test=ast.Compare( + left=ast.Name(id=tmp_name, ctx=ast.Load()), + ops=[ast.IsNot()], + comparators=[ + _create_runtime_attribute( + "_PYIR_SKIP", + lineno=lineno, + col_offset=node.col_offset, + ) + ], + ), + body=[ + ast.Assign( + targets=[_deepcopy_ast_root(target)], + value=ast.Name(id=tmp_name, ctx=ast.Load()), + ) + ], + orelse=[], + ) # target = pyir_read("target", target, owner=..., slot_name=...) # (load before the += computation) target_load_for_read = self._target_as_load(target) pyir_read_call = ast.Call( - func=_create_module_attribute( - "pyir_read", - submodule_name="pyir_runtime", + func=_create_runtime_attribute( + "_pyir_pre_attr_assign" if is_attr_target else "pyir_read", lineno=lineno, col_offset=node.col_offset, ), @@ -2062,12 +5372,25 @@ def _insert_pyir_augassign(self, node: ast.AugAssign) -> list[ast.stmt]: ast.Constant(value=path_str), target_load_for_read, ], - keywords=[deepcopy(kw) for kw in slot_kwargs], - ) - read_stmt = ast.Assign( - targets=[deepcopy(target)], - value=pyir_read_call, + keywords=[_deepcopy_ast_root(kw) for kw in slot_kwargs], ) + if is_attr_target: + read_tmp_name = f"_pyir_read_{self.session_data.counter}" + self.session_data.counter += 1 + read_stmts: list[ast.stmt] = [ + ast.Assign( + targets=[ast.Name(id=read_tmp_name, ctx=ast.Store())], + value=pyir_read_call, + ), + _skip_guarded_store(read_tmp_name), + ] + else: + read_stmts = [ + ast.Assign( + targets=[_deepcopy_ast_root(target)], + value=pyir_read_call, + ) + ] old_name = f"_pyir_old_{self.session_data.counter}" self.session_data.counter += 1 @@ -2083,9 +5406,8 @@ def _insert_pyir_augassign(self, node: ast.AugAssign) -> list[ast.stmt]: # owner=..., slot_name=...) target_load2 = self._target_as_load(target) pyir_call = ast.Call( - func=_create_module_attribute( - "pyir_assign", - submodule_name="pyir_runtime", + func=_create_runtime_attribute( + "_pyir_post_attr_assign" if is_attr_target else "pyir_assign", lineno=lineno, col_offset=node.col_offset, ), @@ -2096,25 +5418,140 @@ def _insert_pyir_augassign(self, node: ast.AugAssign) -> list[ast.stmt]: ast.Constant(value=self.session_data.file_name), ast.Constant(value=lineno), ], - keywords=[deepcopy(kw) for kw in slot_kwargs], - ) - reassign = ast.Assign( - targets=[deepcopy(target)], - value=pyir_call, + keywords=[_deepcopy_ast_root(kw) for kw in slot_kwargs], ) + if is_attr_target: + assign_tmp_name = f"_pyir_assign_{self.session_data.counter}" + self.session_data.counter += 1 + reassign_stmts: list[ast.stmt] = [ + ast.Assign( + targets=[ast.Name(id=assign_tmp_name, ctx=ast.Store())], + value=pyir_call, + ), + _skip_guarded_store(assign_tmp_name), + ] + else: + reassign_stmts = [ + ast.Assign( + targets=[_deepcopy_ast_root(target)], + value=pyir_call, + ) + ] - return [ - ast.copy_location( - ast.fix_missing_locations(read_stmt), node + native_stmts: list[ast.stmt] = [ + *( + ast.copy_location(ast.fix_missing_locations(s), node) + for s in read_stmts ), # load from ref ast.copy_location( ast.fix_missing_locations(capture_old), node ), # capture old - node, # original augmented assignment (uses loaded value) - ast.copy_location( - ast.fix_missing_locations(reassign), node + self._inplace_binop_stmt(node, base_anchor), # boundary-routed in-place op + *( + ast.copy_location(ast.fix_missing_locations(s), node) + for s in reassign_stmts ), # store + reload ] + if not isinstance(target, ast.Attribute): + return native_stmts + + # A property target dispatches through the boundary store (native + # accessor counts: getter once into the in-place op, setter once); + # ordinary attributes keep the storage pattern. + is_prop_test = ast.Call( + func=_create_runtime_attribute( + "_pyir_is_property_store", + lineno=lineno, + col_offset=node.col_offset, + ), + args=[ + self._target_as_load(target.value), + ast.Constant(value=target.attr), + ], + keywords=[], + ) + is_prop_test._pyir_synth = True # type: ignore[attr-defined] + inplace_call = ast.Call( + func=_create_runtime_attribute( + "_pyir_inplace_binop", + lineno=lineno, + col_offset=node.col_offset, + ), + args=[ + ast.Constant(value=self._AUGASSIGN_OP_KEYS[type(node.op)]), + self._target_as_load(target), + _deepcopy_ast_root(node.value), + ], + keywords=[], + ) + prop_store = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_property_store", + lineno=lineno, + col_offset=node.col_offset, + ), + args=[ + self._target_as_load(target.value), + ast.Constant(value=target.attr), + inplace_call, + ], + keywords=[], + ) + ) + if_guard = ast.If( + test=is_prop_test, + body=[ast.copy_location(ast.fix_missing_locations(prop_store), node)], + orelse=native_stmts, + ) + return [ast.copy_location(ast.fix_missing_locations(if_guard), node)] + + def _hoist_subscript_parts( + self, node: ast.stmt, target: ast.Subscript + ) -> list[ast.stmt]: + """Bind an effectful container/key expression of a subscript target + to a temp once and rewrite *target* in place, so the pre-hook, the + guarded writeback, the original statement, and the post-assign all + share the one evaluation (native counts -- a property-backed + container or a call-bearing key otherwise fires per site). + Name/Constant parts stay put: re-reading them is effect-free.""" + hoists: list[ast.stmt] = [] + if not isinstance(target.value, (ast.Name, ast.Constant)): + cont_name = f"_pyir_cont_{self.session_data.counter}" + self.session_data.counter += 1 + cont_hoist = ast.Assign( + targets=[ast.Name(id=cont_name, ctx=ast.Store())], + value=self._target_as_load(target.value), + ) + hoists.append( + ast.copy_location(ast.fix_missing_locations(cont_hoist), node) + ) + target.value = ast.Name(id=cont_name, ctx=ast.Load()) + key = target.slice + if not isinstance(key, (ast.Name, ast.Constant, ast.Slice, ast.Tuple)): + key_name = f"_pyir_key_{self.session_data.counter}" + self.session_data.counter += 1 + key_load = _deepcopy_ast_root(key) + for child in ast.walk(key_load): + if isinstance(child, (ast.Name, ast.Attribute, ast.Subscript)): + child.ctx = ast.Load() + key_hoist = ast.Assign( + targets=[ast.Name(id=key_name, ctx=ast.Store())], + value=key_load, + ) + hoists.append(ast.copy_location(ast.fix_missing_locations(key_hoist), node)) + target.slice = ast.Name(id=key_name, ctx=ast.Load()) + return hoists + + @staticmethod + def _part_hoists_can_effect(hoists: list[ast.stmt]) -> bool: + """True when a hoisted container/key temp's expression can fire an + effect (a call or a walrus binding) if evaluated early.""" + return any( + isinstance(sub, (ast.Call, ast.NamedExpr)) + for h in hoists + for sub in ast.walk(h.value) # type: ignore[attr-defined] + ) def _insert_pyir_subscript_assign( self, node: ast.stmt, target: ast.Subscript @@ -2127,13 +5564,33 @@ def _insert_pyir_subscript_assign( old_name = f"_pyir_sub_old_{self.session_data.counter}" self.session_data.counter += 1 + # One shared evaluation of effectful container/key parts (rewrites + # *target*, and with it the original statement, in place). + part_hoists = self._hoist_subscript_parts(node, target) + # Native assign order is RHS -> container -> key; a part hoist that + # can CALL front-loads its effect ahead of the RHS, so the RHS must + # pin ahead of it. Effect-free hoists keep the plain layout: the + # pre-hook writeback must precede the RHS, whose read of the + # target's own path is un-instrumented (excluded in visit_Assign) + # and reads the container directly. + if self._part_hoists_can_effect(part_hoists) and isinstance(node, ast.Assign): + rhs_name = f"_pyir_rhs_{self.session_data.counter}" + self.session_data.counter += 1 + rhs_hoist = ast.Assign( + targets=[ast.Name(id=rhs_name, ctx=ast.Store())], + value=node.value, + ) + part_hoists.insert( + 0, ast.copy_location(ast.fix_missing_locations(rhs_hoist), node) + ) + node.value = ast.Name(id=rhs_name, ctx=ast.Load()) + container_node = self._target_as_load(target.value) - key_node = deepcopy(target.slice) + key_node = _deepcopy_ast_root(target.slice) pre_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "_pyir_pre_subscript_assign", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), @@ -2152,9 +5609,8 @@ def _insert_pyir_subscript_assign( target_load = self._target_as_load(target) slot_kwargs = self._slot_kwargs_for(target) pyir_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "pyir_assign", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), @@ -2165,16 +5621,15 @@ def _insert_pyir_subscript_assign( ast.Constant(value=self.session_data.file_name), ast.Constant(value=lineno), ], - keywords=[deepcopy(kw) for kw in slot_kwargs], + keywords=[_deepcopy_ast_root(kw) for kw in slot_kwargs], ) reassign = ast.Assign( - targets=[deepcopy(target)], + targets=[_deepcopy_ast_root(target)], value=pyir_call, ) - skip_sentinel = _create_module_attribute( + skip_sentinel = _create_runtime_attribute( "_PYIR_SKIP", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ) @@ -2192,14 +5647,13 @@ def _insert_pyir_subscript_assign( # Writeback: write loaded value back into the dict so the RHS of # the original assignment reads the fresh (ref-loaded) value. # Without this, `d[key] = d[key] + Int32(1)` reads stale d[key]. - skip_sentinel_wb = _create_module_attribute( + skip_sentinel_wb = _create_runtime_attribute( "_PYIR_SKIP", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ) writeback = ast.Assign( - targets=[deepcopy(target)], + targets=[_deepcopy_ast_root(target)], value=ast.Name(id=old_name, ctx=ast.Load()), ) guard_wb = ast.Compare( @@ -2214,6 +5668,7 @@ def _insert_pyir_subscript_assign( ) return [ + *part_hoists, ast.copy_location(ast.fix_missing_locations(pre_assign), node), ast.copy_location(ast.fix_missing_locations(if_writeback), node), node, @@ -2231,13 +5686,16 @@ def _insert_pyir_subscript_augassign( old_name = f"_pyir_sub_old_{self.session_data.counter}" self.session_data.counter += 1 + # One shared evaluation of effectful container/key parts (rewrites + # *target*, and with it the boundary-routed in-place op, in place). + part_hoists = self._hoist_subscript_parts(node, target) + container_node = self._target_as_load(target.value) - key_node = deepcopy(target.slice) + key_node = _deepcopy_ast_root(target.slice) pre_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "_pyir_pre_subscript_assign", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), @@ -2253,14 +5711,13 @@ def _insert_pyir_subscript_augassign( value=pre_call, ) - skip_sentinel_1 = _create_module_attribute( + skip_sentinel_1 = _create_runtime_attribute( "_PYIR_SKIP", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ) writeback = ast.Assign( - targets=[deepcopy(target)], + targets=[_deepcopy_ast_root(target)], value=ast.Name(id=old_name, ctx=ast.Load()), ) guard_1 = ast.Compare( @@ -2277,9 +5734,8 @@ def _insert_pyir_subscript_augassign( target_load = self._target_as_load(target) slot_kwargs = self._slot_kwargs_for(target) pyir_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "pyir_assign", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), @@ -2290,15 +5746,14 @@ def _insert_pyir_subscript_augassign( ast.Constant(value=self.session_data.file_name), ast.Constant(value=lineno), ], - keywords=[deepcopy(kw) for kw in slot_kwargs], + keywords=[_deepcopy_ast_root(kw) for kw in slot_kwargs], ) reassign = ast.Assign( - targets=[deepcopy(target)], + targets=[_deepcopy_ast_root(target)], value=pyir_call, ) - skip_sentinel_2 = _create_module_attribute( + skip_sentinel_2 = _create_runtime_attribute( "_PYIR_SKIP", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ) @@ -2314,21 +5769,54 @@ def _insert_pyir_subscript_augassign( ) return [ + *part_hoists, ast.copy_location(ast.fix_missing_locations(pre_assign), node), ast.copy_location(ast.fix_missing_locations(if_writeback), node), - node, + self._inplace_binop_stmt(node), # boundary-routed in-place op ast.copy_location(ast.fix_missing_locations(if_post), node), ] - def _decompose_tuple_assign( + def _make_rhs_pin(self, node: ast.Assign) -> "tuple[str, ast.stmt]": + """Build ``_pyir_tmp_N = _pyir_pin_tuple_capture()``. + + Evaluates the RHS once, wrapped in the capture pin: Python evaluates + the WHOLE RHS before any element store, so a swap source pins its read + here. Returns the temp name and the statement, so a caller that has to + place the pin itself (a chained assign whose FIRST target is not a + sequence) can emit it ahead of every target. + """ + tmp = f"_pyir_tmp_{self.session_data.counter}" + self.session_data.counter += 1 + pin = ast.Assign( + targets=[ast.Name(id=tmp, ctx=ast.Store())], + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_pin_tuple_capture", + lineno=node.lineno, + col_offset=node.col_offset, + ), + args=[node.value], # the original RHS expression + keywords=[ + ast.keyword(arg="unpack_source", value=ast.Constant(value=True)) + ], + ), + ) + return tmp, ast.copy_location(ast.fix_missing_locations(pin), node) + + def _decompose_unpack_assign( self, node: ast.Assign, - tuple_target: ast.Tuple, + sequence_target: "ast.Tuple | ast.List", in_scope_indices: set[int], *, rhs_temp_name: str | None = None, ) -> ast.stmt | list[ast.stmt]: - """Decompose a tuple unpacking assignment into individual pyir-instrumented assigns. + """Decompose a sequence-unpacking assignment into individual pyir-instrumented assigns. + + ``sequence_target`` is the ``ast.Tuple`` or ``ast.List`` target of the + unpack. Both spellings are accepted because ``[a, b] = rhs`` and + ``a, b = rhs`` compile to identical bytecode -- only the AST node type + differs -- so they must be instrumented the same way. ``in_scope_indices`` is the set of element indices that were in scope BEFORE ``_visit_target`` added first-definitions. This prevents @@ -2351,78 +5839,430 @@ def _decompose_tuple_assign( b = _pyir_tmp_N[1] a = pyir_assign(...) # only if a is in scope b = pyir_assign(...) # only if b is in scope + + A Subscript element ``c[k]`` instead adopts the single-subscript + pre-hook protocol: ``_pyir_old_N = _pyir_pre_subscript_assign(...)`` + classifies the container BEFORE any element access, and the + writeback and ``pyir_assign`` steps are guarded on + ``_pyir_old_N is not _PYIR_SKIP``. """ stmts: list[ast.stmt] = [] lineno = node.lineno col_offset = node.col_offset - elts = tuple_target.elts + elts = sequence_target.elts - # Starred unpacking not supported -- fall through without instrumentation. - if any(isinstance(elt, ast.Starred) for elt in elts): - return node + # Starred unpacking: the star target binds a fresh LIST (never + # instrumented here); positional siblings keep star-aware indices. + star_idx: "int | None" = None + for _i, _elt in enumerate(elts): + if isinstance(_elt, ast.Starred): + if star_idx is not None: + return node # ill-formed (Python allows only one star) + star_idx = _i # Use the pre-computed in-scope indices from visit_Assign. instrumented: list[tuple[int, ast.expr]] = [ (i, elt) for i, elt in enumerate(elts) if i in in_scope_indices ] + # A Name element outside the ACTIVE region scopes is still a first-def + # of the enclosing FUNCTION's place at the unpack position (region + # scopes are an AST artifact; Python scoping is function-wide), so it + # commits through the first-def choke -- the runtime ledger decides + # first-def vs live-row rebind. + first_def_elts: list[tuple[int, ast.Name]] = [ + (i, elt) + for i, elt in enumerate(elts) + if i not in in_scope_indices + and isinstance(elt, ast.Name) + and elt.id != "_" + and not self.session_data.scope_manager.is_skip_reference_taking(elt.id) + ] + + # A possibly-first-def attribute element: gate its read/old-capture/post-assign + # on a ``hasattr`` flag (as scalar assigns do) to avoid ``AttributeError``. + # A property element instead branch-selects on the same value-free + # target fact as scalar assigns: the probe and the read/old/post + # steps would fire the accessors, so the flag stays False and the + # Step-4 store dispatches the bound setter through the boundary. + has_flags: dict[int, str] = {} # index -> _pyir_has_N flag name + prop_flags: dict[int, str] = {} # index -> _pyir_isprop_N flag name + for idx, elt in instrumented: + if isinstance(elt, ast.Attribute): + isprop_name = f"_pyir_isprop_{self.session_data.counter}" + self.session_data.counter += 1 + prop_flags[idx] = isprop_name + is_prop_test = ast.Call( + func=_create_runtime_attribute( + "_pyir_is_property_store", + lineno=lineno, + col_offset=col_offset, + ), + args=[ + self._target_as_load(elt.value), + ast.Constant(value=elt.attr), + ], + keywords=[], + ) + is_prop_test._pyir_synth = True # type: ignore[attr-defined] + isprop_assign = ast.Assign( + targets=[ast.Name(id=isprop_name, ctx=ast.Store())], + value=is_prop_test, + ) + stmts.append( + ast.copy_location(ast.fix_missing_locations(isprop_assign), node) + ) + flag_name = f"_pyir_has_{self.session_data.counter}" + self.session_data.counter += 1 + has_flags[idx] = flag_name + hasattr_call = ast.Call( + func=ast.Name(id="hasattr", ctx=ast.Load()), + args=[ + self._target_as_load(elt.value), + ast.Constant(value=elt.attr), + ], + keywords=[], + ) + # Tracer-internal storage probe, not a user reflection read: + # keep it bare through the call-boundary pass. + hasattr_call._pyir_synth = True # type: ignore[attr-defined] + flag_assign = ast.Assign( + targets=[ast.Name(id=flag_name, ctx=ast.Store())], + value=ast.BoolOp( + op=ast.And(), + values=[ + ast.UnaryOp( + op=ast.Not(), + operand=ast.Name(id=isprop_name, ctx=ast.Load()), + ), + hasattr_call, + ], + ), + ) + stmts.append( + ast.copy_location(ast.fix_missing_locations(flag_assign), node) + ) + + def _maybe_guard(stmt: ast.stmt, idx: int) -> ast.stmt: + """Wrap *stmt* in ``if :`` for possible-first-def attrs.""" + flag = has_flags.get(idx) + if flag is None: + return stmt + guard = ast.If( + test=ast.Name(id=flag, ctx=ast.Load()), + body=[stmt], + orelse=[], + ) + return ast.copy_location(ast.fix_missing_locations(guard), node) + + def _guard_unbound_name(stmt: ast.stmt, on_unbound: list[ast.stmt]) -> ast.stmt: + """Wrap a bare-name element's old-value read/capture so an unbound name + degrades to a first-def instead of crashing.""" + tried = ast.Try( + body=[stmt], + handlers=[ + ast.ExceptHandler( + type=ast.Tuple( + elts=[ + ast.Name(id="NameError", ctx=ast.Load()), + ast.Name(id="UnboundLocalError", ctx=ast.Load()), + ], + ctx=ast.Load(), + ), + name=None, + body=on_unbound or [ast.Pass()], + ) + ], + orelse=[], + finalbody=[], + ) + return ast.copy_location(ast.fix_missing_locations(tried), node) + + def _skip_guarded(stmt: ast.stmt, old_name: str) -> ast.stmt: + """Wrap *stmt* in ``if is not _PYIR_SKIP:`` (the + single-subscript pre-hook guard shape).""" + guard = ast.Compare( + left=ast.Name(id=old_name, ctx=ast.Load()), + ops=[ast.IsNot()], + comparators=[ + _create_runtime_attribute( + "_PYIR_SKIP", + lineno=lineno, + col_offset=col_offset, + ) + ], + ) + guarded = ast.If(test=guard, body=[stmt], orelse=[]) + return ast.copy_location(ast.fix_missing_locations(guarded), node) + + def _pin_rhs() -> str: + tmp, pin = self._make_rhs_pin(node) + stmts.append(pin) + return tmp + + # One shared evaluation of effectful container/key parts per + # subscript element (rewrites each *elt*, and with it the Step-4 + # extraction and Step-5 post-assign, in place); emitted at the + # element's Step-1 slot below. Native unpack evaluates the WHOLE + # RHS before any target part, so a hoist that can CALL pulls the + # Step-3 RHS pin ahead of it; effect-free hoists keep the plain + # layout (their Step-1 writeback must precede the RHS's + # un-instrumented same-path reads). + part_hoists_by_idx = { + idx: self._hoist_subscript_parts(node, elt) + for idx, elt in instrumented + if isinstance(elt, ast.Subscript) + } + if rhs_temp_name is None and any( + self._part_hoists_can_effect(h) for h in part_hoists_by_idx.values() + ): + rhs_temp_name = _pin_rhs() + # --- Step 1: pyir_read for each instrumented element --- - for _i, elt in instrumented: + # Subscript elements adopt the single-subscript pre-hook protocol + # (fusing Step 2): one runtime call classifies the container and + # captures the old value -- or returns ``_PYIR_SKIP`` BEFORE any + # element access -- and a guarded writeback re-binds the element so + # the RHS reads the fresh (ref-loaded) value. + old_names: dict[int, str] = {} # index -> old_name + for idx, elt in instrumented: path_str = self._target_to_path_str(elt) + if isinstance(elt, ast.Subscript): + old_name = f"_pyir_old_{self.session_data.counter}" + self.session_data.counter += 1 + old_names[idx] = old_name + stmts.extend(part_hoists_by_idx[idx]) + pre_call = ast.Call( + func=_create_runtime_attribute( + "_pyir_pre_subscript_assign", + lineno=lineno, + col_offset=col_offset, + ), + args=[ + ast.Constant(value=path_str), + self._target_as_load(elt.value), + _deepcopy_ast_root(elt.slice), + ], + keywords=[], + ) + pre_assign = ast.Assign( + targets=[ast.Name(id=old_name, ctx=ast.Store())], + value=pre_call, + ) + stmts.append( + ast.copy_location(ast.fix_missing_locations(pre_assign), node) + ) + writeback = ast.Assign( + targets=[_deepcopy_ast_root(elt)], + value=ast.Name(id=old_name, ctx=ast.Load()), + ) + writeback = ast.copy_location( + ast.fix_missing_locations(writeback), node + ) + stmts.append(_skip_guarded(writeback, old_name)) + continue target_load = self._target_as_load(elt) slot_kwargs = self._slot_kwargs_for(elt) read_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "pyir_read", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), args=[ast.Constant(value=path_str), target_load], - keywords=[deepcopy(kw) for kw in slot_kwargs], + keywords=[_deepcopy_ast_root(kw) for kw in slot_kwargs], ) - read_stmt = ast.Assign( - targets=[deepcopy(elt)], + read_stmt: ast.stmt = ast.Assign( + targets=[_deepcopy_ast_root(elt)], value=read_call, ) - stmts.append(ast.copy_location(ast.fix_missing_locations(read_stmt), node)) + read_stmt = ast.copy_location(ast.fix_missing_locations(read_stmt), node) + if isinstance(elt, ast.Attribute): + stmts.append(_maybe_guard(read_stmt, idx)) + else: + # Bare name possibly unbound (const_expr-branch first-def): skip + # the read if unbound; Step 4 binds it and Step 5 sees old=None. + stmts.append(_guard_unbound_name(read_stmt, [])) - # --- Step 2: capture old values --- - old_names: dict[int, str] = {} # index -> old_name - for _i, elt in instrumented: + # Step 2: capture old values. For possible-first-def attrs, capture only inside + # the ``hasattr`` guard; the ``else`` arm seeds ``_pyir_old`` with ``None``. + # (Subscript elements were captured by the Step-1 pre-hook.) + for idx, elt in instrumented: + if isinstance(elt, ast.Subscript): + continue old_name = f"_pyir_old_{self.session_data.counter}" self.session_data.counter += 1 - old_names[_i] = old_name + old_names[idx] = old_name target_load = self._target_as_load(elt) - capture = ast.Assign( + capture: ast.stmt = ast.Assign( targets=[ast.Name(id=old_name, ctx=ast.Store())], value=target_load, ) - stmts.append(ast.copy_location(ast.fix_missing_locations(capture), node)) + capture = ast.copy_location(ast.fix_missing_locations(capture), node) + flag = has_flags.get(idx) + if flag is None: + # Bare name may be unbound (const_expr-branch first-def); seed + # _pyir_old = None on an unbound read so Step 5 treats it as a first-def. + name_else_seed = ast.Assign( + targets=[ast.Name(id=old_name, ctx=ast.Store())], + value=ast.Constant(value=None), + ) + stmts.append( + _guard_unbound_name( + capture, + [ + ast.copy_location( + ast.fix_missing_locations(name_else_seed), node + ) + ], + ) + ) + else: + else_seed = ast.Assign( + targets=[ast.Name(id=old_name, ctx=ast.Store())], + value=ast.Constant(value=None), + ) + guard = ast.If( + test=ast.Name(id=flag, ctx=ast.Load()), + body=[capture], + orelse=[ + ast.copy_location(ast.fix_missing_locations(else_seed), node) + ], + ) + stmts.append(ast.copy_location(ast.fix_missing_locations(guard), node)) # --- Step 3: evaluate RHS into temp --- if rhs_temp_name is not None: - # RHS already evaluated by a prior tuple target — reuse the temp. + # RHS already evaluated — by a prior tuple target (multi-target + # assignment) or by the pre-Step-1 pin above. tmp_name = rhs_temp_name else: - tmp_name = f"_pyir_tmp_{self.session_data.counter}" - self.session_data.counter += 1 - tmp_assign = ast.Assign( - targets=[ast.Name(id=tmp_name, ctx=ast.Store())], - value=node.value, # the original RHS expression + tmp_name = _pin_rhs() + + # --- Step 3b: this target's own arity, checked against the temp --- + # Step 4 stores by index, and indexing does not check length: a surplus + # value is dropped silently and a missing one raises ``IndexError`` + # where Python raises ``ValueError``. The check restores both. + # + # Per TARGET, not per statement: targets of one chained assign can want + # different lengths (``(p, q) = (c, r, s) = rhs`` -- Python rejects one + # of them). ``_split_multiple_targets`` has already turned that into + # one store statement per target, so each arrives here on its own and + # validates the shared ``_pyir_rhs_N`` against its own element count. + arity_check = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_check_unpack_arity", + lineno=lineno, + col_offset=col_offset, + ), + args=[ + ast.Name(id=tmp_name, ctx=ast.Load()), + ast.Constant(value=len(elts) - (1 if star_idx is not None else 0)), + ast.Constant(value=star_idx is not None), + ], + keywords=[], ) - stmts.append(ast.copy_location(ast.fix_missing_locations(tmp_assign), node)) + ) + stmts.append(ast.copy_location(ast.fix_missing_locations(arity_check), node)) # --- Step 4: decompose -- extract each element from temp --- + n_elts = len(elts) for i, elt in enumerate(elts): - subscript = ast.Subscript( - value=ast.Name(id=tmp_name, ctx=ast.Load()), - slice=ast.Constant(value=i), - ctx=ast.Load(), - ) + extract_value: ast.expr + if star_idx is not None and i == star_idx: + assert isinstance(elt, ast.Starred) + # rest = list(tmp[star : star - (n-1) or end]) + tail_count = n_elts - 1 - star_idx + slice_node = ast.Slice( + lower=ast.Constant(value=star_idx), + upper=ast.Constant(value=-tail_count) if tail_count else None, + step=None, + ) + extract_value = ast.Call( + func=ast.Name(id="list", ctx=ast.Load()), + args=[ + ast.Subscript( + value=ast.Name(id=tmp_name, ctx=ast.Load()), + slice=slice_node, + ctx=ast.Load(), + ) + ], + keywords=[], + ) + target_node: ast.expr = _deepcopy_ast_root(elt.value) # strip the Star + else: + # Position from the left before the star, from the RIGHT after. + idx_val = i if star_idx is None or i < star_idx else i - n_elts + extract_value = ast.Subscript( + value=ast.Name(id=tmp_name, ctx=ast.Load()), + slice=ast.Constant(value=idx_val), + ctx=ast.Load(), + ) + target_node = _deepcopy_ast_root(elt) + prop_flag = prop_flags.get(i) + if prop_flag is not None and isinstance(elt, ast.Attribute): + # Extract once, then branch the store: a property target runs + # its bound setter through the boundary (observed, predicated); + # an ordinary attribute keeps the native store. + elt_name = f"_pyir_elt_{self.session_data.counter}" + self.session_data.counter += 1 + elt_hoist = ast.Assign( + targets=[ast.Name(id=elt_name, ctx=ast.Store())], + value=extract_value, + ) + stmts.append( + ast.copy_location(ast.fix_missing_locations(elt_hoist), node) + ) + prop_store = ast.Expr( + value=ast.Call( + func=_create_runtime_attribute( + "_pyir_property_store", + lineno=lineno, + col_offset=col_offset, + ), + args=[ + self._target_as_load(elt.value), + ast.Constant(value=elt.attr), + ast.Name(id=elt_name, ctx=ast.Load()), + ], + keywords=[], + ) + ) + plain_store = ast.Assign( + targets=[target_node], + value=ast.Name(id=elt_name, ctx=ast.Load()), + ) + branch = ast.If( + test=ast.Name(id=prop_flag, ctx=ast.Load()), + body=[ + ast.copy_location(ast.fix_missing_locations(prop_store), node) + ], + orelse=[ + ast.copy_location(ast.fix_missing_locations(plain_store), node) + ], + ) + stmts.append(ast.copy_location(ast.fix_missing_locations(branch), node)) + continue + nest_root = elt.value if isinstance(elt, ast.Starred) else elt + if isinstance(nest_root, ast.Name) and nest_root.id.startswith( + "_pyir_nest_" + ): + # Nested-unpack temp (split in visit_Assign): pin its elements + # HERE, before this statement's Step-5 stores, so the follow-up + # unpack consumes capture-time values (swap safety). + extract_value = ast.Call( + func=_create_runtime_attribute( + "_pyir_pin_tuple_capture", + lineno=lineno, + col_offset=col_offset, + ), + args=[extract_value], + keywords=[], + ) extract = ast.Assign( - targets=[deepcopy(elt)], - value=subscript, + targets=[target_node], + value=extract_value, ) stmts.append(ast.copy_location(ast.fix_missing_locations(extract), node)) @@ -2433,9 +6273,8 @@ def _decompose_tuple_assign( slot_kwargs = self._slot_kwargs_for(elt) old_name = old_names[idx] pyir_call = ast.Call( - func=_create_module_attribute( + func=_create_runtime_attribute( "pyir_assign", - submodule_name="pyir_runtime", lineno=lineno, col_offset=col_offset, ), @@ -2446,13 +6285,42 @@ def _decompose_tuple_assign( ast.Constant(value=self.session_data.file_name), ast.Constant(value=lineno), ], - keywords=[deepcopy(kw) for kw in slot_kwargs], + keywords=[_deepcopy_ast_root(kw) for kw in slot_kwargs], ) - reassign = ast.Assign( - targets=[deepcopy(elt)], + reassign: ast.stmt = ast.Assign( + targets=[_deepcopy_ast_root(elt)], value=pyir_call, ) - stmts.append(ast.copy_location(ast.fix_missing_locations(reassign), node)) + reassign = ast.copy_location(ast.fix_missing_locations(reassign), node) + if isinstance(elt, ast.Subscript): + # Pre-hook-skipped elements record no facts and never fire + # another element access (single-subscript guard shape). + stmts.append(_skip_guarded(reassign, old_name)) + else: + stmts.append(_maybe_guard(reassign, idx)) + + # --- Step 5b: first-def commit for out-of-scope Name elements --- + for idx, elt in first_def_elts: + fd_call = ast.Call( + func=_create_runtime_attribute( + "pyir_assign", + lineno=lineno, + col_offset=col_offset, + ), + args=[ + ast.Constant(value=elt.id), + ast.Constant(value=None), + self._target_as_load(elt), + ast.Constant(value=self.session_data.file_name), + ast.Constant(value=lineno), + ], + keywords=[], + ) + fd_stmt: ast.stmt = ast.Assign( + targets=[_deepcopy_ast_root(elt)], + value=fd_call, + ) + stmts.append(ast.copy_location(ast.fix_missing_locations(fd_stmt), node)) if len(stmts) == 0: return node diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_runtime.py b/python/CuTeDSL/cutlass/base_dsl/pyir_runtime.py index 1ce8583cb0..65501740c0 100644 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_runtime.py +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_runtime.py @@ -13,130 +13,18 @@ """PyIR runtime facade -- re-exports the layered pyir_* split. ``pyir_assign`` / ``pyir_read`` / ``MutableValue`` and the full prior -``pyir_runtime`` surface are defined across the pyir_state -> ... -> pyir_cleanup -chain; this module re-exports them so ``pyir_runtime.`` keeps working. +``pyir_runtime`` surface are defined across the pyir_state -> ... -> +pyir_call_boundary chain; this module re-exports them so +``pyir_runtime.`` keeps working. No ``__all__`` here: nothing +star-imports this facade, and named imports resolve through the chain +modules' generated ``__all__`` lists (see scripts/gen_pyir_all.py). """ -from .pyir_cleanup import * # noqa: F401,F403 (top of the layer chain -> pulls in everything) - -# -- BEGIN explicit imports for the type checker (do not edit the list by hand; -# it mirrors the chain modules' module-level definitions, which reach this -# facade at runtime via the wildcard + each layer's dynamic ``__all__`` -- -# constructs a static type checker cannot evaluate. Purely additive: the -# wildcard import above stays the runtime source of truth. -from .pyir_state import ( # noqa: F401 - _meta_uses as _meta_uses, - _slot_refs as _slot_refs, - _slot_first_def_inside_cf as _slot_first_def_inside_cf, - _slot_first_def_depth as _slot_first_def_depth, - _slot_first_def_block as _slot_first_def_block, - _slot_first_def_depth_any as _slot_first_def_depth_any, - _pinned_owners as _pinned_owners, - _slot_mvs as _slot_mvs, - _slot_pending_store as _slot_pending_store, - _pyir_fn_id_stack as _pyir_fn_id_stack, - pyir_function_scope as pyir_function_scope, - _current_fn_id as _current_fn_id, - _SCF_REGION_NAMES as _SCF_REGION_NAMES, - _PYIR_SKIP as _PYIR_SKIP, - _SLOT_STORE_ATTR as _SLOT_STORE_ATTR, - _PYIR_SLOT_FALLBACK as _PYIR_SLOT_FALLBACK, - _PYIR_LIST_MUTATORS as _PYIR_LIST_MUTATORS, - _PYIR_DICT_MUTATORS as _PYIR_DICT_MUTATORS, - _PYIR_SET_MUTATORS as _PYIR_SET_MUTATORS, - _PYIR_READ_SIMPLE_SENTINEL as _PYIR_READ_SIMPLE_SENTINEL, - _NON_DOT_FUNC_ENTRY_OPS as _NON_DOT_FUNC_ENTRY_OPS, - _MODULE_OPS as _MODULE_OPS, -) -from .pyir_core import ( # noqa: F401 - _make_slot_key as _make_slot_key, - _cached_ir_value_dominates_ip as _cached_ir_value_dominates_ip, - _emit_constant_at_ip as _emit_constant_at_ip, - _unwrap as _unwrap, - _wrapper_slots as _wrapper_slots, - _record_fold_witness as _record_fold_witness, - _WatchedM as _WatchedM, - _WatchedInt as _WatchedInt, - _WatchedBool as _WatchedBool, - _WatchedFloat as _WatchedFloat, - _replace_value_uses as _replace_value_uses, - _exit_function_trace as _exit_function_trace, - _watched_to_dsl as _watched_to_dsl, - _describe_value_origin as _describe_value_origin, - _FUNC_OPS as _FUNC_OPS, - _auto_promote_primitive as _auto_promote_primitive, - _can_create_ref as _can_create_ref, - _is_vector_like as _is_vector_like, - _is_memref_like as _is_memref_like, - _mlir_type_or_none as _mlir_type_or_none, - _staged_type_changed as _staged_type_changed, - _is_boolean_like as _is_boolean_like, - _is_literal_backed as _is_literal_backed, - _POISON_EMITTED as _POISON_EMITTED, - _make_poison_like as _make_poison_like, - _record_poison_source as _record_poison_source, - _get_defining_operation as _get_defining_operation, - _get_function_entry_block as _get_function_entry_block, - _create_ref as _create_ref, - _pyir_lookup_slot_from_value as _pyir_lookup_slot_from_value, - _pyir_value_tracked_by_accessible_ref as _pyir_value_tracked_by_accessible_ref, - _load_as_dsl as _load_as_dsl, - _ancestor_op_in_block as _ancestor_op_in_block, - _slot_store_for_tier1 as _slot_store_for_tier1, - _slot_store_for_tier2 as _slot_store_for_tier2, - _slot_storage_available as _slot_storage_available, - _registry_owner as _registry_owner, - _get_slot_mv as _get_slot_mv, - _set_slot_mv as _set_slot_mv, - _clear_slot_mv as _clear_slot_mv, - _iter_slot_mvs_for_pyir_read as _iter_slot_mvs_for_pyir_read, - _attach_mutable_ref as _attach_mutable_ref, - _fresh_wrapper as _fresh_wrapper, - MutableValue as MutableValue, - _get_instance_attrs as _get_instance_attrs, - _is_compound_single_leaf as _is_compound_single_leaf, - _is_leaf_decomposable as _is_leaf_decomposable, - _check_tuple_decomposable as _check_tuple_decomposable, - _check_all_fields_decomposable as _check_all_fields_decomposable, - _flatten_tuple as _flatten_tuple, - _pyir_assign_simple as _pyir_assign_simple, - _is_func_boundary_op as _is_func_boundary_op, - _is_module_boundary_op as _is_module_boundary_op, - _raw_backing_ir_value as _raw_backing_ir_value, - _value_dominates_current_ip as _value_dominates_current_ip, - _same_ir_value as _same_ir_value, - _op_is_inside_op as _op_is_inside_op, - _innermost_enclosing_loop_op_at_ip as _innermost_enclosing_loop_op_at_ip, - _loop_free_enclosing_if_ops_at_ip as _loop_free_enclosing_if_ops_at_ip, -) -from .pyir_corewalk import ( # noqa: F401 - _pyir_auto_load_arg as _pyir_auto_load_arg, - _raw_ir_value as _raw_ir_value, - _arg_value_dominates_current_ip as _arg_value_dominates_current_ip, - _has_decomposable_staged_fields as _has_decomposable_staged_fields, - _has_any_staged_content as _has_any_staged_content, -) -from .pyir_threading import ( # noqa: F401 - _meta_promote_slot as _meta_promote_slot, -) -from .pyir_entrypoints import ( # noqa: F401 - _decompose_tuple as _decompose_tuple, - _decompose_m2m_assign as _decompose_m2m_assign, - _pyir_check_no_complex_m2m_call as _pyir_check_no_complex_m2m_call, - pyir_tag_pending_writes as pyir_tag_pending_writes, - pyir_promote_loop_body_arg as pyir_promote_loop_body_arg, - pyir_assign as pyir_assign, - pyir_read as pyir_read, - _pyir_post_subscript_read as _pyir_post_subscript_read, - _subscript_container_is_dsl_managed as _subscript_container_is_dsl_managed, - _pyir_pre_subscript_assign as _pyir_pre_subscript_assign, - with_ctxmgr_check as with_ctxmgr_check, - _pyir_while_cond as _pyir_while_cond, -) -from .pyir_cleanup import ( # noqa: F401 - _verify_no_used_poison as _verify_no_used_poison, -) -# -- END explicit imports for the type checker - - -__all__ = [name for name in list(globals()) if not name.startswith("__")] +from .pyir_state import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_core import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_corewalk import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_loop_carry import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_entrypoints import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_call_boundary import * # noqa: F401,F403 (siblings: each layer imports ALL lower layers) +from .pyir_class_facts import * # noqa: F401,F403 (class-facts leaf; not part of the layer chain) +from .pyir_spec import * # noqa: F401,F403 (F-SPEC re-entry engine leaf; not part of the layer chain) diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_spec.py b/python/CuTeDSL/cutlass/base_dsl/pyir_spec.py new file mode 100644 index 0000000000..65a97300c7 --- /dev/null +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_spec.py @@ -0,0 +1,464 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: LicenseRef-NvidiaProprietary +# +# Use of this software is governed by the terms and conditions of the +# NVIDIA End User License Agreement (EULA), available at: +# https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/license.html +# +# Any use, reproduction, disclosure, or distribution of this software +# and related documentation outside the scope permitted by the EULA +# is strictly prohibited. + + +"""F-SPEC re-entry verification engine (V-11 SPEC-MATCH): re-resolves every +sealed specialization root path against a live call and refuses on drift. +The executor keeps only the ``_pyir_spec`` attr and a lazy hook into +``_pyir_validate_spec_reentry``.""" + +import sys +import types +from typing import Any, TYPE_CHECKING + +from .common import DSLUserCodeError +from .diagnostics import DiagId +from .pyir_state import ( + _Sentinel, + _SPEC_ATTR_ABSENT, + _SPEC_ATTR_PRESENT, + _SPEC_RECEIVER_DEAD, + _SpecContainerSnapshot, + _SpecTraceExitObject, +) + +if TYPE_CHECKING: + # Type-only: the executor imports this module lazily at validation time, + # so a runtime import here would close an import cycle. + from .jit_executor import ExecutionArgs + + +class _SpecUnresolvable: + """Sentinel: a specialization root path does not resolve on this call.""" + + +_SPEC_UNRESOLVABLE = _SpecUnresolvable() + + +class _SpecUnsupplied: + """Sentinel: the root's argument is not part of THIS call's graph (an + omitted trace-time-constant argument stays the compile-time capture).""" + + +_SPEC_UNSUPPLIED = _SpecUnsupplied() + + +class _SpecBinderFailed: + """Sentinel: the launch binder could not produce this call's binding.""" + + +_SPEC_BINDER_FAILED = _SpecBinderFailed() + + +def _spec_bound_call_arguments( + execution_args: "ExecutionArgs | None", args: tuple, kwargs: dict +) -> "dict[str, Any] | _SpecBinderFailed": + """THIS call's argument binding, produced by the launch binder itself + (one name->position fact source, so validation == launch by construction); + a receiver or constexpr-pruned parameter is never positional-bound.""" + if execution_args is None: + return _SPEC_BINDER_FAILED + try: + return execution_args.bound_call_arguments(args, kwargs) + except DSLUserCodeError: + # The launch's own argument diagnostic: surface it here (fail closed). + raise + except Exception: + return _SPEC_BINDER_FAILED + + +def _spec_resolve_root_path( + root_path: tuple, + bound_args: "dict[str, Any] | None", + entry_func: Any, +) -> Any: + """Resolve one (kind, key, steps) specialization root path against THIS + call's argument graph / the entry function's own live roots.""" + kind, key, steps = root_path + try: + if kind == "arg": + if bound_args is None or key not in bound_args: + # The argument is absent from this call's graph: the handle + # keeps its compile-time capture; nothing to compare. + return _SPEC_UNSUPPLIED + cur = bound_args[key] + elif kind == "global": + fn_globals = getattr(entry_func, "__globals__", None) + if not fn_globals or key not in fn_globals: + return _SPEC_UNRESOLVABLE + cur = fn_globals[key] + elif kind == "module": + cur = sys.modules.get(key) + if cur is None: + return _SPEC_UNRESOLVABLE + elif kind == "closure": + closure = getattr(entry_func, "__closure__", None) + if not closure or key >= len(closure): + return _SPEC_UNRESOLVABLE + cur = closure[key].cell_contents + else: + return _SPEC_UNRESOLVABLE + for step_kind, step_key in steps: + cur = _spec_step_into(cur, step_kind, step_key) + return cur + except DSLUserCodeError: + raise # the purity guard's curated refusal is never demoted + except Exception: + return _SPEC_UNRESOLVABLE + + +# C-level storage descriptors (slot members / getsets) read memory and run no +# user code: STORAGE by declared fact (DF-2). +_SPEC_STORAGE_DESCRIPTORS = (types.MemberDescriptorType, types.GetSetDescriptorType) +_SPEC_STAGE_MISS = _Sentinel("spec stage miss") + + +def _spec_attr_stage(cur: Any, key: str) -> "tuple[str, Any]": + """Classify one attribute hop on *cur* along CPython's own attribute + precedence (data descriptor -> instance storage -> non-data class entry -> + ``__getattr__``), each stage STORAGE or CODE by structural facts alone + (DF-2) -- nothing user-defined executes here. Returns ``("data", value)`` + for a storage answer or ``("code", stage_kind)`` when reaching the value + would execute class code; raises ``AttributeError`` on structural absence. + A type-level ``__getattribute__`` override changes native semantics + invisibly to a structural walk, so every hop on such a type is CODE. + Class and module owners are symbol NAMESPACES -- never wrapper-family -- + so their native lookup protocol is the declared (DF-5) channel.""" + if isinstance(cur, (type, types.ModuleType)): + return ("code", "namespace-lookup") + klass = type(cur) + getattribute = getattr(klass, "__getattribute__", object.__getattribute__) + if ( + getattribute is not object.__getattribute__ + and getattribute is not type.__getattribute__ + ): + return ("code", "__getattribute__") + class_entry: Any = _SPEC_STAGE_MISS + for k in klass.__mro__: + if key in k.__dict__: + class_entry = k.__dict__[key] + break + if class_entry is not _SPEC_STAGE_MISS: + entry_type = type(class_entry) + if isinstance(class_entry, _SPEC_STORAGE_DESCRIPTORS): + # Declared C-storage read; an unset slot falls through to the + # __getattr__ stage exactly as native lookup does. + try: + return ("data", class_entry.__get__(cur, klass)) + except AttributeError: + class_entry = _SPEC_STAGE_MISS + elif hasattr(entry_type, "__get__") and ( + hasattr(entry_type, "__set__") or hasattr(entry_type, "__delete__") + ): + return ("code", "data-descriptor") + inst = getattr(cur, "__dict__", None) + if isinstance(inst, dict): + val = dict.get(inst, key, _SPEC_STAGE_MISS) + if val is not _SPEC_STAGE_MISS: + return ("data", val) + if class_entry is not _SPEC_STAGE_MISS: + if hasattr(type(class_entry), "__get__"): + return ("code", "non-data-descriptor") + return ("data", class_entry) + if getattr(klass, "__getattr__", None) is not None: + return ("code", "__getattr__") + raise AttributeError(key) + + +def _spec_step_into(cur: Any, step_kind: str, step_key: Any) -> Any: + """Re-resolve one recorded step against a live object. An attr step is + an owner-SLOT hop: container adoption models dict items as owner slots, + so a mapping at the hop resolves its slot by key; any unknown step kind + or failing access raises (the caller reads it as unresolvable). On an + MLIR value-tree WRAPPER (its class implements ``__extract_mlir_values__``) + an attr hop resolves structurally only: executing ANY code stage there + (descriptor or ``__getattr__``) would re-run trace-time IR construction + against consumed SSA (a crash channel), so it refuses loudly as a + record-time classification bug. + A plain live object keeps full attribute semantics -- re-running its + property against live state is the designed exact re-derivation.""" + if step_kind == "attr": + if isinstance(cur, dict): + return cur[step_key] + stage, payload = _spec_attr_stage(cur, str(step_key)) + if stage == "data": + return payload + if getattr(type(cur), "__extract_mlir_values__", None) is not None: + raise DSLUserCodeError( + DiagId.SPEC_DESCRIPTOR_HOP, + attr=str(step_key), + owner_type=type(cur).__name__, + path=f"{type(cur).__name__}.{step_key}", + ) + # DF-5: the declared plain-object code channel -- re-running a live + # property/__getattr__ against live state is the exact re-derivation. + return getattr(cur, step_key) + if step_kind == "item": + return cur[step_key] + if step_kind == "cell": + return cur.__closure__[step_key].cell_contents + if step_kind == "default": + return cur.__defaults__[step_key] + if step_kind == "kwdefault": + return cur.__kwdefaults__[step_key] + raise KeyError(step_kind) + + +def _spec_scalarize(value: Any) -> Any: + """Reduce a value-semantic numeric wrapper to its constexpr Python scalar so + a baked scalar row compares BY VALUE against a live wrapper (a fresh Int32 + carrying the same number is not stale); any other value passes through + unchanged (opaque objects keep identity).""" + try: + from .pyir_core import _NO_CONST_VALUE, _pyir_meta_primitive_value + + prim = _pyir_meta_primitive_value(value) + return prim if prim is not _NO_CONST_VALUE else value + except Exception: + return value + + +def _spec_payloads_equal(baked: Any, live: Any) -> bool: + """Exact structural equality between a baked payload and a live value + (bool never equates to int at any depth -- tuple interiors compare + element-wise; a failing/ambiguous ``==`` is a mismatch). A trace-written + object row verifies IDENTITY: the live value must BE the trace-exit + object. A live value-semantic numeric wrapper reduces to its scalar first, + so a baked scalar verifies its VALUE, not the fresh wrapper's identity.""" + if isinstance(baked, _SpecTraceExitObject): + return baked.matches(live) + if isinstance(baked, _SpecContainerSnapshot): + return baked.matches(live) + live = _spec_scalarize(live) + baked = _spec_scalarize(baked) + if isinstance(baked, bool) != isinstance(live, bool): + return False + if isinstance(baked, tuple) or isinstance(live, tuple): + if type(baked) is not type(live) or len(baked) != len(live): + return False + return all(_spec_payloads_equal(b, v) for b, v in zip(baked, live)) + try: + return (baked == live) is True + except Exception: + return False + + +def _spec_format_root_path(root_path: tuple) -> str: + """Author-readable spelling of a specialization root path.""" + kind, key, steps = root_path + if kind == "arg": + text = f"argument {key}" if isinstance(key, int) else f"argument '{key}'" + elif kind == "global": + text = f"global '{key}'" + elif kind == "module": + text = f"module '{key}'" + else: + text = f"closure cell {key}" + for step_kind, step_key in steps: + if step_kind == "attr": + text += f".{step_key}" + elif step_kind == "item": + text += f"[{step_key!r}]" + elif step_kind == "default": + text += f".__defaults__[{step_key}]" + elif step_kind == "kwdefault": + text += f".__kwdefaults__[{step_key!r}]" + else: + text += f"." + return text + + +def _spec_resolve_unsupplied_root( + root_path: tuple, entry_func: Any, receiver: Any, sig_facts: "tuple | None" +) -> Any: + """An argument absent from the launch binding is a compile-time capture + living on the handle itself. Its live verification root is the sealed + RECEIVER (a bound method's fixed first binding) or the entry function's + CURRENT default object -- both can drift after compile, so both are read + live and compared. Call structure comes from the SEALED signature facts + (DF-1); the default read is ``__defaults__[i]`` / ``__kwdefaults__[name]`` + -- pure data, never signature-code execution. A root with neither + channel keeps the compile-time capture only when nothing at the call + could name it; a receiver-rooted row with no sealed receiver, or a + default-rooted row whose index map sealed inconsistent, is unverifiable.""" + kind, key, steps = root_path + if kind != "arg": + return _SPEC_UNSUPPLIED + if sig_facts is None: + return _SPEC_UNSUPPLIED + first_param, default_index = sig_facts + root: Any = _SPEC_UNSUPPLIED + if key == first_param and receiver is not None: + if receiver is _SPEC_RECEIVER_DEAD: + # The record roots at the receiver parameter but the receiver + # died before the seal: nothing live can re-verify the row. + return _SPEC_UNRESOLVABLE + root = receiver + if root is _SPEC_UNSUPPLIED and default_index is not None: + fact = default_index.get(key) + if fact is not None: + if fact[0] == "pos": + dflts = getattr(entry_func, "__defaults__", None) or () + if fact[1] >= len(dflts): + return _SPEC_UNRESOLVABLE # defaults shape drifted + root = dflts[fact[1]] + else: + kwd = getattr(entry_func, "__kwdefaults__", None) or {} + if key not in kwd: + return _SPEC_UNRESOLVABLE # kw-only default dropped + root = kwd[key] + elif root is _SPEC_UNSUPPLIED and default_index is None: + # Seal-time signature/live-shape disagreement: nothing here can be + # index-resolved -- the row is unverifiable (fail closed, never + # index-shifted). + return _SPEC_UNRESOLVABLE + if root is _SPEC_UNSUPPLIED: + return _SPEC_UNSUPPLIED + cur = root + try: + for step_kind, step_key in steps: + cur = _spec_step_into(cur, step_kind, step_key) + return cur + except DSLUserCodeError: + raise # the purity guard's curated refusal is never demoted + except Exception: + return _SPEC_UNRESOLVABLE + + +def _spec_verify_probe_row( + root_path: tuple, + baked: Any, + bound_args: "dict[str, Any] | None", + entry_func: Any, + receiver: Any, + sig_facts: "tuple | None", +) -> None: + """Re-run a presence/absence PROBE row: the owner prefix must still + resolve (fail closed otherwise) and the final attr hop must re-answer the + baked boolean -- a flipped probe answer refuses instead of re-entering + stale; the value behind a present probe is free to change.""" + kind, key, steps = root_path + if not steps: + raise DSLUserCodeError(DiagId.SPEC_UNVERIFIABLE_REENTRY) + prefix_path = (kind, key, steps[:-1]) + owner = _spec_resolve_root_path(prefix_path, bound_args, entry_func) + if owner is _SPEC_UNSUPPLIED: + owner = _spec_resolve_unsupplied_root( + prefix_path, entry_func, receiver, sig_facts + ) + if owner is _SPEC_UNSUPPLIED: + return # omitted-argument root: the handle keeps its capture + if owner is _SPEC_UNRESOLVABLE: + raise DSLUserCodeError( + DiagId.STALE_SPECIALIZATION, + path=_spec_format_root_path(root_path), + baked=repr(baked), + live="", + ) + try: + _spec_step_into(owner, *steps[-1]) + live_present = True + except (AttributeError, KeyError, IndexError, TypeError): + live_present = False + if live_present == (baked is _SPEC_ATTR_PRESENT): + return + raise DSLUserCodeError( + DiagId.STALE_SPECIALIZATION, + path=_spec_format_root_path(root_path), + baked=repr(baked), + live="" % ("present" if live_present else "absent"), + ) + + +def _pyir_validate_spec_reentry( + spec: "tuple | None", + execution_args: "ExecutionArgs | None", + args: tuple, + kwargs: dict, +) -> None: + """V-11 SPEC-MATCH at the only no-retrace reuse event: every recorded root + path must re-resolve to its baked exact payload, and the record must be + complete -- a compiled handle cannot re-trace, so a mismatch refuses. + Rows the launch binding does not name verify through the handle's own + live captures (sealed receiver, signature defaults) instead of being + silently skipped.""" + if spec is None: + return + record, complete, entry_func, receiver = spec[:4] + sig_facts = spec[4] if len(spec) > 4 else None + bound_args = ( + _spec_bound_call_arguments(execution_args, args, kwargs) if record else None + ) + if isinstance(bound_args, _SpecBinderFailed): + # An arg-rooted record cannot be verified without the launch binding; + # a compiled handle cannot re-trace, so fail closed. + if any(rp[0] == "arg" for rp in record): + raise DSLUserCodeError(DiagId.SPEC_UNVERIFIABLE_REENTRY) + bound_args = None + for root_path, baked in record.items(): + if baked is _SPEC_ATTR_ABSENT or baked is _SPEC_ATTR_PRESENT: + _spec_verify_probe_row( + root_path, baked, bound_args, entry_func, receiver, sig_facts + ) + continue + live = _spec_resolve_root_path(root_path, bound_args, entry_func) + if live is _SPEC_UNSUPPLIED: + live = _spec_resolve_unsupplied_root( + root_path, entry_func, receiver, sig_facts + ) + if live is _SPEC_UNSUPPLIED: + continue + if live is _SPEC_UNRESOLVABLE or not _spec_payloads_equal(baked, live): + raise DSLUserCodeError( + DiagId.STALE_SPECIALIZATION, + path=_spec_format_root_path(root_path), + baked=repr(baked), + live="" if live is _SPEC_UNRESOLVABLE else repr(live), + ) + if not complete: + raise DSLUserCodeError(DiagId.SPEC_UNVERIFIABLE_REENTRY) + + +# Static export surface (regenerate with scripts/gen_pyir_all.py): the +# module's OWN names -- its namespace minus what lower chain modules +# already export and minus the generator's INTERNAL_ONLY registry. +__all__ = [ + "sys", + "types", + "Any", + "TYPE_CHECKING", + "DSLUserCodeError", + "DiagId", + "_Sentinel", + "_SPEC_ATTR_ABSENT", + "_SPEC_ATTR_PRESENT", + "_SPEC_RECEIVER_DEAD", + "_SpecContainerSnapshot", + "_SpecTraceExitObject", + "_SpecUnresolvable", + "_SPEC_UNRESOLVABLE", + "_SpecUnsupplied", + "_SPEC_UNSUPPLIED", + "_SpecBinderFailed", + "_SPEC_BINDER_FAILED", + "_spec_bound_call_arguments", + "_spec_resolve_root_path", + "_SPEC_STORAGE_DESCRIPTORS", + "_SPEC_STAGE_MISS", + "_spec_attr_stage", + "_spec_step_into", + "_spec_scalarize", + "_spec_payloads_equal", + "_spec_format_root_path", + "_spec_resolve_unsupplied_root", + "_spec_verify_probe_row", + "_pyir_validate_spec_reentry", +] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_state.py b/python/CuTeDSL/cutlass/base_dsl/pyir_state.py index 407c2c1801..72e8710724 100644 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_state.py +++ b/python/CuTeDSL/cutlass/base_dsl/pyir_state.py @@ -12,32 +12,44 @@ """PyIR runtime -- state layer; see facade for the public surface.""" -# ============================================================================= # Standard library imports -# ============================================================================= import inspect +import sys +import types +import weakref as _weakref -from collections.abc import Iterator -from contextlib import contextmanager -from typing import TYPE_CHECKING, Any, Optional -from weakref import WeakKeyDictionary +from typing import ( + TYPE_CHECKING, + Any, + Callable, + Literal as _Literal, + NamedTuple as _NamedTuple, + Optional, +) + +if TYPE_CHECKING: + # Type-only: a runtime import of pyir_core would close an import cycle, so + # these names appear only inside string annotations. + from .pyir_core import MutableValue, _SlotId -# ============================================================================= # Local imports -# ============================================================================= from .multi_stage_manager import ( is_inside_staged_cf, is_inside_locally_staged_cf, is_inside_constexpr_loop, - get_staged_cf_depth, + constexpr_scope_under_staged_cf, + current_staged_cf_depth, _is_staged_value, assign_meta_staged_check, ) -# ============================================================================= # MLIR imports -# ============================================================================= -from .common import DSLRuntimeError, DSLUserCodeError, is_auto_m2s_enabled +from .common import ( + DSLRuntimeError, + DSLUserCodeError, + is_auto_m2s_enabled, + is_pyir_enabled, +) from .diagnostics import DiagId from .._mlir import ir from .utils.logger import log @@ -51,231 +63,434 @@ except ImportError: ub = None -if TYPE_CHECKING: - # Annotation-only; a runtime import would cycle (pyir_core imports this). - from .pyir_core import MutableValue -# Unbind the guard so the layer chain's dynamic ``__all__`` surface is unchanged. -del TYPE_CHECKING - -# --------------------------------------------------------------------------- -# D1 (META_VALUE_TABLE_DESIGN) state and types — retroactive promotion of M values -# --------------------------------------------------------------------------- -# -# When a Python (M) value is read inside staged control flow it is baked -# into the IR as ``arith.constant`` via the DSL value-coercion path. If -# the same slot is later written, we must retroactively replace each -# baked constant with ``pyir.load %ref`` so the IR observes the -# mutation. The two tables below are the bookkeeping for this: -# -# _meta_uses[slot_key] -> list of arith.constant ops baked from -# leaf reads of that slot. Only LEAF -# constants are tracked; derived arith ops -# live downstream via SSA edges and get -# rewritten automatically when the leaf is -# replaced. -# _slot_refs[slot_key] -> pyir.ref handle once the slot has been -# promoted. Subsequent reads return a -# fresh pyir.load wrapped in the matching -# DSL Numeric. -# -# Both dicts are per-function-trace and cleared at trace exit via -# ``_exit_function_trace()`` (driven from ``multi_stage_manager._jit_scope``). -_meta_uses: "dict[Any, list[ir.Value]]" = {} +# PyIR-mode fact per DSL env prefix, registered when the env manager is created. +# Half-enabled tracing is a supported mode for helper jits (one DSL layer may +# legitimately trace with PyIR while a helper layer traces without), so no +# refusal keys on this; the verify-failure funnel consults it to NAME the +# mixed-flag configuration when a cross-DSL kernel compile leaves +# region-crossing SSA behind (an otherwise-opaque dominance ICE). +_PYIR_MODE_BY_PREFIX: "dict[str, bool]" = {} -_slot_refs: "dict[Any, ir.Value]" = {} +def _pyir_register_mode_fact(prefix: str, enabled: bool) -> None: + """Record one DSL prefix's PyIR enablement (declared config fact), and bind + the Numeric wrapper write funnel the first time any prefix enables PyIR.""" + _PYIR_MODE_BY_PREFIX[str(prefix)] = bool(enabled) + if enabled: + from .typing import _pyir_install_numeric_write_funnel + _pyir_install_numeric_write_funnel() -# Per-slot flag set at first-def time -- True when the slot's first -# assignment ran with ``is_inside_staged_cf() == True``. Consulted by -# ``pyir_assign``'s D1 reassignment branch to refuse promotion when the -# user's logical semantics is "reset every iteration" (the slot was born -# inside a runtime-iterated region with a Python-literal init). -_slot_first_def_inside_cf: "dict[Any, bool]" = {} +def _pyir_mixed_mode_hint() -> "str | None": + """A configuration sentence when the process's DSL prefixes disagree on + PyIR enablement, or None when the modes are coherent.""" + if len(set(_PYIR_MODE_BY_PREFIX.values())) < 2: + return None + on = sorted(p for p, v in _PYIR_MODE_BY_PREFIX.items() if v) + off = sorted(p for p, v in _PYIR_MODE_BY_PREFIX.items() if not v) + return ( + "PyIR is enabled for {on} but not for {off}: a kernel compiled " + "across the two DSL layers mixes traced modes, which can leave " + "region-crossing SSA behind (the verification failure above). Set " + "the *_ENABLE_PYIR flags coherently for every DSL layer in the " + "trace.".format(on=", ".join(on), off=", ".join(off)) + ) -# Staged-CF nesting depth recorded at each slot's first-def (parallel to -# ``_slot_first_def_inside_cf``). The keep-constexpr gate keeps a -# Python-primitive toggle constexpr ONLY when the reassignment happens at -# the SAME depth as the first-def. A strictly greater depth means the -# toggle is inside a NESTED RUNTIME loop/if entered after the first-def -# (e.g. ``scale_d`` first-def'd in an outer ``cutlass.range`` then toggled -# inside a nested ``range`` MMA-K loop) -- there the value MUST be promoted -# to a ``pyir.ref`` so the runtime loop threads it as an iter_arg. Keeping -# it constexpr drops the iter_arg and the op reads the stale init forever -# (L103: FMHA-decode QK GEMM accumulate flag stuck at ``false``). -_slot_first_def_depth: "dict[Any, int]" = {} + +# --- Dominance-check helpers (used by _mlir_helpers/op.py) --- + +_SCF_REGION_NAMES = { + "scf.for": "a for-loop body", + "scf.if": "an if/else body", + "scf.while": "a while-loop body", +} -# MLIR block in which each slot was first-defined (parallel to the dicts -# above). When a deferred promotion later inserts a per-iteration reset -# store (see ``_meta_promote_slot``), the reset must land in the block -# where the user wrote the first-def -- NOT before the earliest baked use, -# which may sit inside a NESTED runtime loop entered after the first-def. -# Resetting there would re-reset the slot on every inner iteration (L103). -_slot_first_def_block: "dict[Any, Any]" = {} +# Ref placement helpers +# Function-like entry ops NOT following the ``*.func`` convention; ``*.func`` ops +# are matched structurally by :func:`_is_func_boundary_op`. +_NON_DOT_FUNC_ENTRY_OPS = frozenset(("cuda.kernel",)) -# Staged-CF nesting depth recorded at EVERY local's first-def, regardless of -# value kind (primitive OR staged DSL value). Distinct from -# ``_slot_first_def_depth`` above, which is recorded only for Python-primitive -# inits and feeds the D1 keep-constexpr gate -- overloading it would entangle -# that gate with this discriminator. -# -# Consulted by ``pyir_assign``'s straight-line type-change rebind guard: a -# same-name local reassigned to a value of a DIFFERENT MLIR type is a Rule-4 -# "unstable join" violation ONLY when the prior binding crosses a staged-CF -# boundary (the prior binding pre-dates the current region, so the type would -# differ at a branch merge / loop back-edge). When the prior binding was -# first-defined at the SAME depth as the reassignment, the two assignments are -# straight-line code (e.g. ``p = base + off`` then ``p = inttoptr(p)``); the -# type change is sequential dataflow, not a join, and is handled by the -# existing type-changing-reassignment path (a fresh ref typed to the new -# value). Non-PyIR accepts this; PyIR must too. +# Symbol-table container ops: the entry-block walk terminates here (no SSA +# region to host a ``pyir.ref``). +_MODULE_OPS = frozenset(("builtin.module", "gpu.module")) + + +# Retroactive promotion: ``_meta_uses[slot]`` = LEAF constants baked from reads in +# staged CF; ``_slot_refs[slot]`` = promoted ref. Per-trace, cleared at exit. + +_meta_uses: "dict[Any, list[ir.Value]]" = {} +# Idempotent Mp->Mp rebinds skipped by the promotion gates leave an anchor +# constant at each write site. If the slot LATER promotes (a different-value +# write), ``_meta_promote_slot`` re-materialises every skipped write as a +# per-site reset store so the cell is reset exactly where Python executed +# the assignment. Never-promoted anchors are dead constants (canonicalized +# away). Per-trace, cleared at exit. +_meta_idempotent_write_anchors: "dict[Any, list[ir.Value]]" = {} +_slot_refs: "dict[Any, ir.Value]" = {} +# Wrapper-template half of the bare-ref row above (F-TYPEID): the wrapper the +# place's last store/publish choke held, advanced per store; reads reconstruct +# the row through it (never from a read-site sample). Per-trace, cleared at exit. +_slot_templates: "dict[Any, Any]" = {} +# Gate-once evidence for trace-folded ``if`` gates: the compare choke fills the +# register only for a provable ``not in`` over a plain set; the folder consumes it. +_PYIR_LAST_NOTIN_COMPARE: "list[Any]" = [None] +_PYIR_FOLD_FIRSTDEF_STACK: "list[list[Any]]" = [] +# Tag for attr first-def entries on the fold stack (plain slot keys are locals). +CF_ATTR_FIRST_DEF = "cf-attr-first-def" +# Slots whose first-def was proven a once-per-key init (latched membership gate): +# a later meta arithmetic write is the loop-carried advance and promotes. +_PYIR_GATE_ONCE_INIT_SLOTS: "set[Any]" = set() +# Depth of plain-callee executions observed by a call boundary: first-defs there +# have unprovable multiplicity, so keep-meta must not apply. +_PYIR_BOUNDARY_CALLEE_DEPTH: "list[int]" = [0] + +# True when the slot's first assignment ran inside staged CF. +_slot_first_def_inside_cf: "dict[Any, bool]" = {} +# Birth block of a place's first binding (F-BIRTHPOS): the insertion block open +# when the first-def choke ran. A region-born cell seeds itself with a real +# store at this block instead of a fabricated dominating init. +_slot_first_def_block: "dict[Any, ir.Block]" = {} +# Birth block recorded for EVERY local first-def (F-BIRTHPOS, meta bindings +# too): a rebind in the SAME block is provably trace-local -- no back-edge or +# branch merge separates birth from rebind. +_slot_first_def_block_any: "dict[Any, ir.Block]" = {} +# Staged-CF depth at the slot's literal first-def: distinguishes an outer-loop +# per-iteration reset from a same-level accumulator. +_slot_first_def_depth: "dict[Any, int]" = {} +# Like ``_slot_first_def_depth`` but recorded for EVERY local first-def (compound +# staged values too); the type-change-at-join guard compares it to the current depth. _slot_first_def_depth_any: "dict[Any, int]" = {} +# Staged-CF depth at which the slot's CURRENT binding was last written; the +# type-stability gate distinguishes a same-depth shadow from a shallower loop-carry. +_slot_binding_depth: "dict[Any, int]" = {} +# Innermost-open staged-loop body blocks (F-BIRTHPOS declaration): a first-def +# choke running while a body block is open records that block as the binding's +# birth position. +_pyir_open_loop_body_blocks: "list[ir.Block]" = [] -# A slot key embeds id(owner), and CPython reuses an object's address once it -# is freed. If an owner is dropped mid-trace, a later object can land on the -# same address, build the identical slot key, and alias the dead owner's stale -# entries -- a wrong-result hazard. Hold a strong reference to every owner we -# key on so it cannot be freed (and its address cannot be reused) while the -# trace is live. Keyed by id(owner) so re-keying the same live owner overwrites -# a single entry. Released in _exit_function_trace() when the trace ends. -_pinned_owners: "dict[int, Any]" = {} - - -# Slot-keyed MutableValue registry for owners that cannot host per-object slot -# storage (dict / list / non-__dict__ / non-weakref-able owners). Keyed by the -# same ``_make_slot_key`` tuple used for ``_slot_refs`` / ``_meta_uses`` -- the -# owner is pinned via ``_pinned_owners`` so id(owner) stays stable, so this needs -# no separate finalizer. __dict__-backed objects use per-object tier-1/tier-2 -# storage instead and never land here. Value type is ``MutableValue`` (Any here -# to avoid a circular import with pyir_runtime). Cleared at trace exit. -_slot_mvs: "dict[Any, Any]" = {} - -# Fold-witness table (PHASE_PREDICATE_FOLDED_STALE): slot_key -> list of -# witness records. A witness is recorded when a ``_WatchedM`` fold is -# consumed by PLAIN-CPython control flow inside staged CF -- i.e. the folded -# Python value decided a trace-time branch that no later D1 promotion can -# rewrite (no SSA edge exists from a taken CPython branch to any tracked -# leaf). Two events record a witness: -# -# * ``_WatchedM.__bool__`` -- a truth-test (``if w:``, ``while w:``, -# ``not w``, ``w and x``) of a slot-connected (directly or via -# ``_origin_slots`` provenance) wrapper inside staged CF, and -# * ``_cmp_ir``'s plain-Python fallback returns -- the comparison folded to -# a bare ``bool`` (no derived wrapper with live IR), so any downstream -# consumption is invisible to the meta table. -# -# The witness alone is NOT an error: a fold on a slot that stays constexpr -# for the whole trace is exactly the supported meta-programming model. The -# trace only becomes unsound when the SAME slot is later promoted to a -# ``pyir.ref`` (its reads become runtime loads while the witnessed branch -# stays hard-wired to the stale trace-time value), so ``_meta_promote_slot`` -# checks this table and raises ``DiagId.PHASE_PREDICATE_FOLDED_STALE`` citing -# both the fold site and the mutation site. Cleared at trace exit. -# -# Witness record keys: ``filename``/``lineno`` (fold site, via -# ``find_user_source_location``), ``value`` (folded Python value), ``kind`` -# (human-readable consumption kind). -_fold_witnesses: "dict[Any, list[dict[str, Any]]]" = {} - -# The diagnostic cites only the first fold site plus a count, so any cap >= 2 -# is behaviorally identical; this value merely bounds per-slot memory. -_MAX_FOLD_WITNESSES_PER_SLOT = 8 - -# Slots in the SYNTACTIC write-set of an open staged region (tagged by the -# before-block / body-entry prologue): a genuine store WILL land this trace -# even if it hasn't yet. Lets the while-condition lift stage `n > 0` when -# `n` is loop-carried, while a never-stored predicate keeps its trace-time -# fold (the keep-constexpr contract). Cleared at trace exit. -_slot_pending_store: "set[Any]" = set() - - -# Stack of USER ``@cute.jit`` function ids. ``pyir_function_scope`` pushes one -# for the enclosed user function body. Preprocessor-generated -# loop-body functions do NOT push, so they inherit their enclosing user -# function's id. This makes a frame-local slot key unique per USER function (so -# same-named locals in different functions -- e.g. ``stage_idx`` in copy_sfa and -# copy_sfb -- get DISTINCT slots) while keeping a loop-carried local (which the -# preprocessor threads through generated loop-body functions) on ONE key. -_pyir_fn_id_stack: "list[Any]" = [] +# Trace-global region-epoch clock: bumped at every staged-region entry AND exit +# (depth cannot serve: an enter+exit round trip returns to the same depth), so +# two choke events with equal epochs lie in one straight-line trace interval. +_PYIR_REGION_EPOCH: "list[int]" = [0] +# Region-ENTRY clock + stack of entry indices for the currently-open staged +# regions: a value choke-stamped at entry-clock C is re-readable through its +# cell only where every open region predates C (no re-executing region was +# entered after the stamp, so no later-traced store can precede the read). +_PYIR_REGION_ENTRY_CLOCK: "list[int]" = [0] +_PYIR_REGION_ENTRY_STACK: "list[int]" = [] -@contextmanager -def pyir_function_scope(fn_id: Any) -> Iterator[None]: - """Context manager for the current user-function scope.""" - _pyir_fn_id_stack.append(fn_id) - try: - yield - finally: - _pyir_fn_id_stack.pop() +def _pyir_bump_region_epoch() -> None: + """Advance the region-epoch clock: a staged-region boundary is crossed.""" + _PYIR_REGION_EPOCH[0] += 1 -def _current_fn_id() -> Any: - """Id of the user function currently being traced (top of the stack), or None - outside any tracked function.""" - return _pyir_fn_id_stack[-1] if _pyir_fn_id_stack else None +def _pyir_current_region_epoch() -> int: + """The current region-epoch (recorded on cells at store/load chokes).""" + return _PYIR_REGION_EPOCH[0] -# --------------------------------------------------------------------------- -# Dominance-check helpers (used by _mlir_helpers/op.py) -# --------------------------------------------------------------------------- -_SCF_REGION_NAMES = { - "scf.for": "a for-loop body", - "scf.if": "an if/else body", - "scf.while": "a while-loop body", -} +def _pyir_region_entry_push() -> None: + """Record a staged-region ENTRY on the entry clock and the open stack.""" + _PYIR_REGION_ENTRY_CLOCK[0] += 1 + _PYIR_REGION_ENTRY_STACK.append(_PYIR_REGION_ENTRY_CLOCK[0]) -# ========================================================================== -# Runtime-classified subscript instrumentation -# ========================================================================== +def _pyir_region_entry_pop() -> None: + """Close the innermost recorded staged-region entry (no-op when empty).""" + if _PYIR_REGION_ENTRY_STACK: + _PYIR_REGION_ENTRY_STACK.pop() -_PYIR_SKIP = object() # Sentinel: container is not a dict, skip instrumentation +def _pyir_region_entry_clock() -> int: + """The current region-entry clock value (stamped on choke products).""" + return _PYIR_REGION_ENTRY_CLOCK[0] -# --------------------------------------------------------------------------- -# Slot-keyed ref identity (slot-keyed MutableValue registry) -# --------------------------------------------------------------------------- -# -# PyIR historically attached ``_mutable_ref`` to the value object. When the -# same Python value is bound into multiple storage slots (three attributes -# share a seed, a dict has three entries with the same value), they collapse -# onto a single ``pyir.ref``. Pure Python rebinds by slot, not by value -# identity, so the tracking must do the same. -# -# Tier 1: owner.__dict__["__pyir_slots__"] -- a dict per owner that maps -# slot_name -> MutableValue. ``object.__setattr__`` bypasses -# ``@dataclass(frozen=True)``. -# -# Tier 2: ``_PYIR_SLOT_FALLBACK`` -- module-level ``WeakKeyDictionary`` that -# maps owner instances (declaring ``__slots__`` with no ``__dict__``) to -# their per-owner ``{slot_name: MutableValue}`` map. -# -# Tier 3: return ``None``. ``pyir_assign``/``pyir_read`` fall back to the -# legacy value-keyed ``_mutable_ref`` attribute on the value. -# -# Built-in ``dict`` / ``list`` owners can carry neither ``__pyir_slots__`` -# (no writable ``__dict__``) nor a weakref. ``_set_slot_mv`` silently skips -# recording for those owners, so subscript lookups fall through to tier 3 -# exactly as before. + +def _pyir_region_entry_watermark() -> int: + """Entry index of the innermost open staged region (0 when none).""" + return _PYIR_REGION_ENTRY_STACK[-1] if _PYIR_REGION_ENTRY_STACK else 0 + + +# Unified slot registry keyed by the structural triple ``_SlotId(kind, id(owner), key)``: +# refs follow storage location, not the object in the slot. Weakref finalizer purges on GC. +_SLOT_REGISTRY: "dict[_SlotId, MutableValue]" = {} +_OWNER_KEEPALIVE: "dict[int, _weakref.ReferenceType]" = {} + +# Owner places promoted through the D1 write choke: id(owner) -> (owner handle, +# keys). Feeds the stale-leaf region reload; trace-scoped, cleared at exit. +_PYIR_PROMOTED_PLACE_LEAVES: "dict[int, tuple[Any, set]]" = {} + + +class _Sentinel: + """A named absence/miss marker: identity-compared like a bare ``object()`` + but self-describing in reprs and debugger views.""" + + __slots__ = ("_name",) + + def __init__(self, name: str) -> None: + self._name = name + + def __repr__(self) -> str: + return f"<{self._name}>" + + +_NO_CONST_VALUE = _Sentinel("no const value") + + +# Watched-container adoption (dict entries / list elements as places) + +# Adopt-once registry: ``id(raw dict|list) -> watched instance`` so aliases converge +# on one object and one set of places; trace-scoped, raw containers kept alive. +_WATCHED_CONTAINER_ADOPTIONS: "dict[int, Any]" = {} +_WATCHED_CONTAINER_KEEPALIVE: "list[Any]" = [] +# F-SPEC root chaining: ``twin token -> (holder token, hop step)`` declared at +# the adoption choke, so container-leg places resolve to re-resolvable root +# paths; trace-scoped (cleared with the adoption registry, NOT at ledger open, +# because argument adoption precedes the F-SPEC open for the same trace). +_PYIR_SPEC_CONTAINER_CHAIN: "dict[int, tuple[int, tuple]]" = {} +# Memo of object ids the adoption walk already visited this trace, so the +# per-sighting walk at the candidate-holder choke is amortised to one traversal. +# Visited objects are pinned for the trace so a recycled id can never be +# mistaken for an already-walked object. +_WATCHED_CONTAINER_WALKED: "set[int]" = set() +_WATCHED_CONTAINER_WALKED_KEEPALIVE: "list[Any]" = [] +# Re-entrancy depth: nested watched dunders raw-delegate while a choke handles an +# access; traces are single-carried, so a depth counter suffices. +_WATCHED_CONTAINER_BYPASS: "list[int]" = [0] +# Choke hooks, installed by the entrypoints layer at import (the hook bodies +# need pyir_read / pyir_assign, which live layers above this module). +_WATCHED_DICT_READ_HOOK: "list[Any]" = [None] +_WATCHED_DICT_WRITE_HOOK: "list[Any]" = [None] +_WATCHED_DICT_GET_MISS_HOOK: "list[Any]" = [None] +_WATCHED_DICT_MUTATOR_HOOK: "list[Any]" = [None] +_WATCHED_DICT_ITER_HOOK: "list[Any]" = [None] +# List-leg hooks: an integer-index leg guarded by the length/order structural +# rule -- separate hook bodies, shared adoption tables. +_WATCHED_LIST_READ_HOOK: "list[Any]" = [None] +_WATCHED_LIST_WRITE_HOOK: "list[Any]" = [None] +_WATCHED_LIST_MUTATOR_HOOK: "list[Any]" = [None] + +# ``id(obj)`` membership set for objects held as tracked-container leg values: +# their attributes are places one leg deeper; trace-scoped. +_WATCHED_CONTAINER_HELD_OBJECTS: "set[int]" = set() +_WATCHED_CONTAINER_HELD_KEEPALIVE: "list[Any]" = [] + +# Registry of the trace's keepalive containers: bookkeeping references that +# pin traced values against GC. Referrer audits (e.g. the fresh-leaf descend +# witness) must treat these as non-aliases. Every keepalive list MUST +# register here at its definition site -- an unregistered keepalive makes +# referrer audits conservatively read its bookkeeping reference as a user +# alias, silently degrading the analyses built on them. +_PYIR_TRACE_KEEPALIVES: "list[Any]" = [ + _WATCHED_CONTAINER_KEEPALIVE, + _WATCHED_CONTAINER_WALKED_KEEPALIVE, + _WATCHED_CONTAINER_HELD_KEEPALIVE, +] + + +# Track refs by storage SLOT, not value: ``__dict__``-backed owners store under +# ``__pyir_slots__`` (tier-1); other owners route through ``_SLOT_REGISTRY``. _SLOT_STORE_ATTR = "__pyir_slots__" +# Every object that registered a slot this trace, keyed by ``id(owner)`` in sighting +# order (plain dict: deterministic iteration; never invokes user __hash__/__eq__). +_PYIR_SLOT_HOLDERS: "dict[int, Any]" = {} -_PYIR_SLOT_FALLBACK: "WeakKeyDictionary[Any, dict[Any, MutableValue]]" = ( - WeakKeyDictionary() -) + +# Trace-scoped registry of objects sighted at PyIR choke points that can root a +# captured-holder walk; same keying/weakref rules as ``_PYIR_SLOT_HOLDERS``. +_PYIR_CANDIDATE_HOLDERS: "dict[int, Any]" = {} + +# Sighting front gate: ids of REGISTERED candidates whose repeat sighting is +# provably a no-op (adoption memoized or protocol-filtered), skipped before +# the adoption walk. Strictly a subset of the registry's ids: an id drops +# with its registry row and the set clears with the registry at trace exit. +_PYIR_SIGHTING_SETTLED: "set[int]" = set() + +# Fast gate for candidate registration: armed at trace-arg intake, disarmed at +# outermost trace exit, so non-PyIR sighting sites pay one truth test. +_PYIR_CANDIDATE_REGISTRY_ACTIVE: "list[bool]" = [False] + + +# Holder storage write clock: one monotone counter; every declared storage +# write on a registered holder stamps the holder row with the current tick. +_PYIR_WRITE_CLOCK: "list[int]" = [0] + +# ``id(holder) -> last-write tick`` for registered weakrefable candidates; +# rows are minted at registration and die with the registry entry, so a row +# never outlives (or predates) the object it describes. +_PYIR_HOLDER_WRITE_STAMPS: "dict[int, int]" = {} + +# ``id(holder) -> (stamp_at_build, records tuple, guards tuple)`` gather-walk +# segments for funneled flat roots (builtin scalars, bare ``ir.Value``, +# SSA-backed Numeric leaves, walk-inert values); root stamp AND per-record +# wrapper guard stamps are validated at splice time; rows are dropped with the +# registry entry (trace-scoped). +_PYIR_GATHER_SEGMENTS: "dict[int, tuple]" = {} + +# The declared wrapper write funnel: ``(base class, payload property, +# __setattr__, __delattr__)`` registered by the class that owns the funnel. +_PYIR_WRITE_FUNNEL_DECL: "list[Any]" = [None] + +# Per-class memo of the write-funnel fact (class shape is definition-stable). +_PYIR_WRITE_FUNNEL_CLASS_FACTS: "dict[type, bool]" = {} + +# Per-class ``{attr: bool}`` memo: ``getattr(instance, attr)`` resolves to the +# instance-storage cell (no foreign MRO data descriptor shadows *attr*). +_PYIR_GETATTR_STORAGE_FACTS: "dict[type, dict]" = {} + + +def _pyir_note_holder_write(obj: Any) -> None: + """Stamp a declared storage write on *obj*'s holder row (no-op when the + object is not a registered candidate).""" + oid = id(obj) + if oid in _PYIR_HOLDER_WRITE_STAMPS: + _PYIR_WRITE_CLOCK[0] += 1 + _PYIR_HOLDER_WRITE_STAMPS[oid] = _PYIR_WRITE_CLOCK[0] + + +def _pyir_setattr_raw(obj: Any, name: str, value: Any) -> None: + """The raw-write funnel: ``object.__setattr__`` plus the write stamp, so a + funnel-bypassing internal write stays visible to the storage-epoch facts.""" + object.__setattr__(obj, name, value) + oid = id(obj) + if oid in _PYIR_HOLDER_WRITE_STAMPS: + _PYIR_WRITE_CLOCK[0] += 1 + _PYIR_HOLDER_WRITE_STAMPS[oid] = _PYIR_WRITE_CLOCK[0] + + +def _pyir_delattr_raw(obj: Any, name: str) -> None: + """The raw-delete funnel: ``object.__delattr__`` plus the write stamp.""" + object.__delattr__(obj, name) + oid = id(obj) + if oid in _PYIR_HOLDER_WRITE_STAMPS: + _PYIR_WRITE_CLOCK[0] += 1 + _PYIR_HOLDER_WRITE_STAMPS[oid] = _PYIR_WRITE_CLOCK[0] + + +def _pyir_declare_write_funnel(base_cls: type) -> None: + """Declare *base_cls* as the wrapper write-funnel owner: its verbatim + ``value`` property, ``__setattr__`` and ``__delattr__`` become the identity + anchors of the per-class funnel fact.""" + _PYIR_WRITE_FUNNEL_DECL[0] = ( + base_cls, + base_cls.__dict__["value"], + base_cls.__dict__["__setattr__"], + base_cls.__dict__["__delattr__"], + ) + _PYIR_WRITE_FUNNEL_CLASS_FACTS.clear() + _PYIR_GETATTR_STORAGE_FACTS.clear() -# Mutator methods on built-in containers that we forbid inside staged CF -# when the container is a plain Python ``list``/``dict``/``set`` -- the -# compiler traces the body once and would silently discard the per- -# iteration mutations. +def _pyir_wrapper_write_funneled(cls: type) -> bool: + """DECLARED write-funnel fact: every language-level attribute write or + delete on a *cls* instance lands in plain ``__dict__`` storage AND stamps + the holder write clock. Requires the declared funnel base in the MRO with + its verbatim ``__setattr__``/``__delattr__``/``value`` property resolved + (no override), no custom ``__getattribute__``, and no ``__slots__``.""" + fact = _PYIR_WRITE_FUNNEL_CLASS_FACTS.get(cls) + if fact is not None: + return fact + decl = _PYIR_WRITE_FUNNEL_DECL[0] + ok = False + if decl is not None and isinstance(cls, type): + base, value_prop, setattr_fn, delattr_fn = decl + try: + ok = ( + issubclass(cls, base) + and cls.__setattr__ is setattr_fn + and cls.__delattr__ is delattr_fn + and cls.__getattribute__ is object.__getattribute__ + and all("__slots__" not in k.__dict__ for k in cls.__mro__) + and next( + (k.__dict__["value"] for k in cls.__mro__ if "value" in k.__dict__), + None, + ) + is value_prop + ) + except Exception: + ok = False + _PYIR_WRITE_FUNNEL_CLASS_FACTS[cls] = ok + return ok + + +def _pyir_getattr_reads_storage(cls: type, attr: str) -> bool: + """True when ``getattr(instance, attr)`` returns the instance-storage cell + verbatim on funneled *cls*: either no MRO entry shadows *attr* (instance + dict wins over non-data descriptors), or the entry IS the declared payload + property, whose getter returns the stored cell.""" + per_cls = _PYIR_GETATTR_STORAGE_FACTS.get(cls) + if per_cls is None: + per_cls = _PYIR_GETATTR_STORAGE_FACTS[cls] = {} + fact = per_cls.get(attr) + if fact is not None: + return fact + ok = _pyir_wrapper_write_funneled(cls) + if ok: + decl = _PYIR_WRITE_FUNNEL_DECL[0] + for k in cls.__mro__: + entry = k.__dict__.get(attr) + if entry is None: + continue + if entry is decl[1]: + break + entry_cls = type(entry) + if hasattr(entry_cls, "__set__") or hasattr(entry_cls, "__delete__"): + ok = False # foreign data descriptor: getattr resolves through it + break + per_cls[attr] = ok + return ok + + +def _pyir_plain_storage_setattr(cls: type) -> bool: + """True when *cls* writes attributes through plain object storage: no + override at all, or the declared wrapper funnel (a verbatim delegate whose + only addition is the write stamp).""" + setattr_fn = cls.__setattr__ + if setattr_fn is object.__setattr__: + return True + decl = _PYIR_WRITE_FUNNEL_DECL[0] + return decl is not None and setattr_fn is decl[2] + + +# Inert attr placeholder for a SELF-LEAF snapshot record (the leaf IS the loop-carried +# local); the record category is recognised structurally by the dedicated holder type. +_PYIR_SELF_LEAF_ATTR = "" + + +# Numeric-leaf classification: carryable scalar Numeric wrapper backing kinds. +# not a carryable scalar Numeric wrapper +_NUMERIC_LEAF_NONE: "_Literal['none']" = "none" +# carryable scalar Numeric backed by a baked SSA +_NUMERIC_LEAF_SSA: "_Literal['ssa']" = "ssa" +# carryable scalar Numeric still a meta literal +_NUMERIC_LEAF_META: "_Literal['meta']" = "meta" + + +# Unit attribute marking a carry ``pyir.ref`` so the while-path dedup recognises a +# post-close load of a nested loop's ref (an op attr, not an id set: wrappers re-mint). +_PYIR_LOOP_ITER_ARGS_ATTR = "pyir.loop_iter_args" + + +# Inert attr placeholder for a TOP-LEVEL-TUPLE leaf snapshot record: the tuple itself is +# the holder (recognised structurally by ``_pyir_is_carryable_tuple``); the leaf walk +_PYIR_TUPLE_LEAF_ATTR = "" + + +# Memref write/alias facts are dialect-specific and DECLARED at the IR layer; +# consumers without the pyir binding degrade to no memref-write detection. + + +# Mutator methods forbidden on plain Python containers inside staged CF: the body +# is traced once, so per-iteration mutations would be silently discarded. _PYIR_LIST_MUTATORS = frozenset( { "append", @@ -288,8 +503,6 @@ def _current_fn_id() -> Any: "reverse", } ) - - _PYIR_DICT_MUTATORS = frozenset( { "update", @@ -299,13 +512,12 @@ def _current_fn_id() -> Any: "clear", } ) - - _PYIR_SET_MUTATORS = frozenset( { "add", "discard", "remove", + "pop", "clear", "update", "intersection_update", @@ -313,14 +525,666 @@ def _current_fn_id() -> Any: "symmetric_difference_update", } ) +_PYIR_DEQUE_MUTATORS = frozenset( + { + "append", + "appendleft", + "pop", + "popleft", + "rotate", + "extend", + "extendleft", + "clear", + "remove", + "insert", + "reverse", + } +) -_PYIR_READ_SIMPLE_SENTINEL = object() +# Mutators that store the caller-supplied object itself, so a staged value with a +# live foreign slot binding would be trapped; non-inserting mutators are excluded. +_PYIR_LIST_INSERT_MUTATORS = frozenset({"append", "extend", "insert"}) +_PYIR_DICT_INSERT_MUTATORS = frozenset({"update", "setdefault"}) -_NON_DOT_FUNC_ENTRY_OPS = frozenset(("cuda.kernel",)) +# --- Place-identity layer: the ledger key space --- +# A "place" is a value-independent key for a logical storage location (never an +# id()), so a reconstructed object at the same binding resolves to the SAME key. + +import itertools as _itertools + +# Monotonic id sources; never reset per trace, so a stale cross-trace reference +# can never collide with a fresh id. +_SCOPE_ID_COUNTER = _itertools.count() +_TOKEN_COUNTER = _itertools.count() + + +class _ScopeFrame(_NamedTuple): + scope_id: int + kind: str # 'fn' (user-function/kernel trace) | 'region' (scf region body) + + +# Authoritative scope stack for bare-name identity (no frame walking): one 'fn' +# frame per instrumented activation; a 'region' frame inherits its 'fn' scope_id. +_PYIR_SCOPE_STACK: "list[_ScopeFrame]" = [] + + +class _PlaceSeg: + """One composite place segment: a base slot plus the exact item-hop keys + that reached the leaf (``t[2][0]`` == base ``'t'``, keys ``(2, 0)``). + + Born at the runtime decomposition sites that hold the parent path and the + hop key as separate Python values, so recording them is a projection of + local knowledge, never an inference; the bracketed string spelling is + DISPLAY-ONLY (``__str__``) and is never parsed back. Deliberately neither + a ``str`` nor a ``tuple`` subclass: an unaudited boundary must fail + loudly on it, never absorb it as an attr name or an exact item key.""" + + __slots__ = ("base", "keys") + base: Any + keys: "tuple[Any, ...]" + + def __init__(self, base: Any, keys: "tuple[Any, ...]" = ()) -> None: + _pyir_setattr_raw(self, "base", base) + _pyir_setattr_raw(self, "keys", tuple(keys)) + + def __setattr__(self, name: str, value: Any) -> None: + raise AttributeError("_PlaceSeg is immutable") + + def child(self, key: Any) -> "_PlaceSeg": + """Extend by one item hop (O(1) at each decomposition level).""" + return _PlaceSeg(self.base, self.keys + (key,)) + + def steps(self) -> "tuple[tuple, ...]": + """Mechanical unfold into F-SPEC re-resolution steps; a str base is an + attribute hop, any other base is an exact item key on the owner.""" + first = ( + ("attr", self.base) if isinstance(self.base, str) else ("item", self.base) + ) + return (first, *(("item", k) for k in self.keys)) + + def __eq__(self, other: Any) -> bool: + return ( + type(other) is _PlaceSeg + and other.base == self.base + and other.keys == self.keys + ) + + def __ne__(self, other: Any) -> bool: + return not self.__eq__(other) + + def __hash__(self) -> int: + return hash((_PlaceSeg, self.base, self.keys)) + + def __str__(self) -> str: + # Rendering parity with the historical composite spelling + # (f"{base}[{k!r}]..."): for int keys f"{i}" == f"{i!r}". + return f"{self.base}" + "".join(f"[{k!r}]" for k in self.keys) + + def __repr__(self) -> str: + return f"_PlaceSeg({self.base!r}, {self.keys!r})" + + +def _place_seg_child(parent: Any, key: Any) -> _PlaceSeg: + """One composite hop below *parent*: a ``_PlaceSeg`` parent extends; any + other parent (a plain attr name or an exact container key) is the base.""" + if isinstance(parent, _PlaceSeg): + return parent.child(key) + return _PlaceSeg(parent, (key,)) + + +class IdentityKeyedWeakTable: + """Object->value map keyed by referent IDENTITY (the standard identity-map + algorithm): bookkeeping never invokes user ``__hash__``/``__eq__``, so a + lookup on a staged value can emit no IR and value-equal distinct owners + can never alias one row. Weakrefable keys purge their row at collection + (the callback runs during deallocation, before the address can recycle); + non-weakrefable keys are pinned until :meth:`clear`.""" + + __slots__ = ("_rows",) + + def __init__(self) -> None: + # id(key) -> (anchor, value); anchor is a weakref or the pinned key. + self._rows: "dict[int, tuple[Any, Any]]" = {} + + def get(self, obj: Any, default: Any = None) -> Any: + row = self._rows.get(id(obj)) + return default if row is None else row[1] + + def __contains__(self, obj: Any) -> bool: + return id(obj) in self._rows + + def __setitem__(self, obj: Any, value: Any) -> None: + oid = id(obj) + row = self._rows.get(oid) + if row is not None: + self._rows[oid] = (row[0], value) + return + rows = self._rows + + def _drop_row(_r: Any, _oid: int = oid) -> None: + rows.pop(_oid, None) + + try: + anchor: Any = _weakref.ref(obj, _drop_row) + except TypeError: + anchor = obj + rows[oid] = (anchor, value) + + def pop(self, obj: Any, default: Any = None) -> Any: + row = self._rows.pop(id(obj), None) + return default if row is None else row[1] + + def is_pinned(self, obj: Any) -> bool: + """True when *obj*'s row holds a strong pin (non-weakrefable key).""" + row = self._rows.get(id(obj)) + return row is not None and row[0] is obj + + def clear(self) -> None: + self._rows.clear() + + def __len__(self) -> int: + return len(self._rows) + + +# Stable id()-independent owner tokens, propagated across reconstruction. +# Identity-keyed: token bookkeeping can never route through a payload dunder. +_OWNER_TOKENS: "IdentityKeyedWeakTable" = IdentityKeyedWeakTable() + +# Rebuild-protocol frame bit (F-BIRTH): non-zero while the +# ``__new_from_mlir_values__`` machinery rebuilds a carried object; owner-token +# adoption asserts it (V-6) -- a constructor call births a fresh token. +_PYIR_REBUILD_PROTOCOL_DEPTH: "list[int]" = [0] + +# Owner-token birth class (F-CLASS): token -> type(owner) at mint; carried +# uses validate the live class against it (V-7). +_PYIR_TOKEN_BORN_CLASS: "dict[int, type]" = {} + +# The place-keyed cell store; holds MLIR-context-bound MutableValues, cleared per trace. +_PLACE_REGISTRY: "dict[Any, MutableValue]" = {} + +# Place -> row store-version at the restructuring rebind that could NOT advance +# the row's cell: serving the row unadvanced would read the superseded +# generation, so the place-routed read refuses instead (discharged by any +# later tracked store or a fresh cell registered at the place). +_PYIR_SUPERSEDED_PLACE_ROWS: "dict[Any, int]" = {} + +# Attr-chain places root at the nearest STABLE anchor with a symbolic suffix. +# ``_OWNER_PLACE_PREFIX``: owner object -> (root_token, attr-name suffix from anchor). +_OWNER_PLACE_PREFIX: "IdentityKeyedWeakTable" = IdentityKeyedWeakTable() +# ``_ROOT_NAME_TOKENS``: (scope_id, name) of a dotted access's outermost binding +# name -> anchor token; scoped per activation (a name rebinds per method). +_ROOT_NAME_TOKENS: "dict[tuple, int]" = {} + +# ``(scope_id, name) -> owner tokens of the closure CELLS registered for the +# binding``, in registration order (a frame's entry registers first; region- +# synthesized frames sharing the scope append their carry-param cells). +# A cell is the one binding a nested closure and its owning frame share, so +# the tokens are the exact aliasing facts between a boundary closure read and +# a later local write. +_PYIR_SCOPE_CELL_TOKENS: "dict[tuple[int, str], tuple[int, ...]]" = {} + +# ``cell token -> (scope_id, name) of the cell's OWNING binding`` (the scope +# whose entry registered the cell as its own cellvar, not as a ``nonlocal`` +# freevar); every scope sharing the cell resolves its place key to this row. +_PYIR_CELL_HOME_BINDING: "dict[int, tuple[int, str]]" = {} + +# ``scope_id -> frozenset of names the function DECLARES nonlocal``, registered +# at instrumented-function entry (rewrite-time syntactic fact); a bare-name +# choke consults the innermost scope's set to recognize closure-cell bindings. +_PYIR_SCOPE_NONLOCAL_NAMES: "dict[int, frozenset]" = {} + +# F-CEPLACE binding-birth ownership: ``(scope_id, name) -> serial of the +# constexpr instance that was innermost-open at the binding's BIRTH`` (or +# ``None`` for an outside-born binding; parameters are seeded ``None`` at +# scope entry so the fact is declared, never defaulted). A row persists past +# its instance's close -- the qualified place key stays derivable for post-if +# and post-loop reads -- until a dead-owner rebind births a new binding. +_CE_BINDING_OWNER: "dict[tuple[int, str], int | None]" = {} + +# One-cell-per-place violation records: a staged access whose authoritative and +# place-cell resolutions name two live cells. Appended before the raise; capped. +_EMISSION_SELF_CHECK_LOG: "list[dict]" = [] +_EMISSION_SELF_CHECK_LOG_CAP = 10000 + +# --- F-SPEC: per-trace specialization ledger ------------------------------- +# spec = { root_path -> exact payload } where root_path = (kind, key, steps), +# kind in {"arg", "global", "closure"}; steps are ("attr", name) | ("item", k) +# | ("cell", i) hops from the root object. Payloads are exact values compared +# structurally at executor re-entry -- never repr digests. + +# Owner token -> root path of the object it was minted for (arg graph at trace +# intake, plus lazily identity-resolved globals of the entry function). +_PYIR_SPEC_TOKEN_ROOTS: "dict[int, tuple]" = {} +# Local place ("local", scope_id, name) -> root path (entry parameter locals). +_PYIR_SPEC_LOCAL_ROOTS: "dict[Any, tuple]" = {} +# Entry parameter roots staged at trace intake; adopted by the next 'fn' scope +# (the entry body opens it immediately after intake). +_PYIR_SPEC_PENDING_PARAM_ROOTS: "list[tuple[str, tuple]]" = [] +# Tokens that failed a global-identity scan; never re-scanned this trace. +_PYIR_SPEC_UNROOTED_TOKENS: "set[int]" = set() +# Tokens of TRACE-EPOCH objects (IR wrappers written to persistent places +# during this trace, plus values hopped off them): a row chain passing +# through one is trace-internal -- nothing at a later launch can re-derive +# it, so it never records a launch-recheckable row. +_PYIR_SPEC_TRACE_BORN_TOKENS: "set[int]" = set() +# Strong pins for stamped trace-epoch values: a token in the set provably +# names the stamped object for the whole trace (no dead-object aliasing). +_PYIR_SPEC_TRACE_BORN_PINS: "list[Any]" = [] +# The recorded specialization set (first read per root path wins: the payload +# closest to entry state is what the trace baked). +_PYIR_SPEC_RECORD: "dict[tuple, Any]" = {} +# Cleared by any bake event not pathable to a re-resolvable root (R6b) and by +# a recording failure (an unrecorded bake is unverifiable); an incomplete +# record refuses executor re-entry instead of guessing. +_PYIR_SPEC_COMPLETE: "list[bool]" = [True] +# The traced entry function: the live root for global/closure re-dereference. +_PYIR_SPEC_ENTRY_FUNC: "list[Any]" = [None] +# Entry-signature facts captured at trace OPEN and sealed with the record +# (DF-1): (first_param_name, {name: ("pos", default_idx) | ("kwonly",)}); +# the default map is None when the live __defaults__/__kwdefaults__ shape +# disagreed with the signature at capture (those rows seal unverifiable). +_PYIR_SPEC_ENTRY_SIG: "list[Any]" = [None] +# The traced entry receiver (a bound method's fixed first binding): the live +# root for verifying receiver-rooted rows at re-entry, where the launch +# binding never names the receiver. +_PYIR_SPEC_ENTRY_RECEIVER: "list[Any]" = [None] + + +class _SpecReceiverDead: + """Sentinel sealed when the record roots rows at the receiver parameter + but the receiver object died before the seal: the rows exist and nothing + live can re-verify them (distinct from ``None`` = no receiver was ever + part of this entry, e.g. a plain function).""" + + +_SPEC_RECEIVER_DEAD = _SpecReceiverDead() + + +class _SpecTraceExitObject: + """F-SPEC payload for a row the trace itself WROTE with an object-shaped + value (a receiver field populated during setup, e.g. ``None`` -> struct): + the trace-exit fact is the written object itself, so re-entry verifies + live-value IDENTITY. The referent is held weakly when the object supports + it (the record never extends a lifetime the verified live state does not + already hold) and directly otherwise (tuple subclasses, builtin + containers); a collected referent can no longer be the live value, so it + reads as drift and refuses.""" + + __slots__ = ("_weak", "_strong", "_desc") + + def __init__(self, obj: Any) -> None: + self._desc = f"{type(obj).__module__}.{type(obj).__qualname__}" + try: + self._weak: Any = _weakref.ref(obj) + self._strong: Any = None + except TypeError: + self._weak = None + self._strong = obj + + def matches(self, live: Any) -> bool: + obj = self._strong if self._weak is None else self._weak() + return obj is not None and obj is live + + def __repr__(self) -> str: + return f"" + + +class _SpecContainerSnapshot: + """F-SPEC payload for a WHOLE-container consumption (``x in lst`` routes + through C-level ``__contains__``, so no per-item choke fires): the bake + depends on the full contents, so the row holds an exact contents snapshot + -- scalar legs by value, object legs by identity -- and re-entry compares + length plus every leg against the live container.""" + + __slots__ = ("kind", "elems") + + def __init__(self, kind: str, elems: tuple) -> None: + self.kind = kind # "list" (items) | "dict-keys" (key order + values) + self.elems = elems + + def matches(self, live: Any) -> bool: + # One comparator for every F-SPEC row shape: per-leg equality is + # _spec_payloads_equal (late import: pyir_spec imports this module). + # The retired private _leg_eq had drifted from it -- it scalarized + # only the live side, only via the ``python_value`` protocol, and + # refused nested container-snapshot legs a fresh equal container + # should satisfy. + try: + from .pyir_spec import _spec_payloads_equal + + if self.kind == "list": + if not isinstance(live, list) or list.__len__(live) != len(self.elems): + return False + items: Any = list.__iter__(live) + else: + if not isinstance(live, dict) or dict.__len__(live) != len(self.elems): + return False + items = dict.keys(live) + return all(_spec_payloads_equal(b, v) for b, v in zip(self.elems, items)) + except Exception: + return False + + def __repr__(self) -> str: + return f"<{self.kind} contents snapshot, {len(self.elems)} leg(s)>" + + +class _SpecAttrAbsent: + """F-SPEC payload for a tolerated attribute-probe MISS (``hasattr`` False, + 3-arg ``getattr`` default): the bake depends on the name staying absent, + so re-entry re-probes the final attr hop and refuses when it resolves.""" + + def __repr__(self) -> str: + return "" + + +_SPEC_ATTR_ABSENT = _SpecAttrAbsent() + + +class _SpecAttrPresent: + """F-SPEC payload for a ``hasattr`` storage HIT: the bake is the PRESENCE + boolean (never the value), so re-entry re-probes the final attr hop and + refuses only when the name stops resolving.""" + + def __repr__(self) -> str: + return "" + + +_SPEC_ATTR_PRESENT = _SpecAttrPresent() + + +# Sealed (record, complete, entry_func) of the last closed outer trace; the +# compile flow moves it onto the JitCompiledFunction it builds. +_PYIR_SPEC_SEALED: "list[Any]" = [None] + + +def _reset_spec_state() -> None: + """Clear the live F-SPEC ledger (the sealed snapshot is left in place for + the compile flow to collect).""" + _PYIR_SPEC_TOKEN_ROOTS.clear() + _PYIR_SPEC_LOCAL_ROOTS.clear() + _PYIR_SPEC_PENDING_PARAM_ROOTS.clear() + _PYIR_SPEC_UNROOTED_TOKENS.clear() + _PYIR_SPEC_TRACE_BORN_TOKENS.clear() + _PYIR_SPEC_TRACE_BORN_PINS.clear() + _PYIR_SPEC_RECORD.clear() + _PYIR_SPEC_COMPLETE[0] = True + _PYIR_SPEC_ENTRY_FUNC[0] = None + _PYIR_SPEC_ENTRY_SIG[0] = None + _PYIR_SPEC_ENTRY_RECEIVER[0] = None + + +def _reset_scope_state() -> None: + """Clear per-trace place-layer state; monotonic counters intentionally + survive trace exit.""" + _PYIR_SCOPE_STACK.clear() + _PYIR_REBUILD_PROTOCOL_DEPTH[0] = 0 + _PYIR_TOKEN_BORN_CLASS.clear() + _PLACE_REGISTRY.clear() + _PYIR_SUPERSEDED_PLACE_ROWS.clear() + _ROOT_NAME_TOKENS.clear() + _PYIR_SCOPE_CELL_TOKENS.clear() + _PYIR_CELL_HOME_BINDING.clear() + _PYIR_SCOPE_NONLOCAL_NAMES.clear() + _CE_BINDING_OWNER.clear() + _PYIR_REGION_EPOCH[0] = 0 + _PYIR_REGION_ENTRY_CLOCK[0] = 0 + _PYIR_REGION_ENTRY_STACK.clear() + _OWNER_TOKENS.clear() + _OWNER_PLACE_PREFIX.clear() + _PYIR_GATE_ONCE_INIT_SLOTS.clear() + + +# Base-form spine state: the fn-id stack the entrypoints spine drives; the +# enter/exit hooks double as the ledger scope bracket. + +# Stack of USER jit-decorated function ids (generated loop bodies do not push): +# frame-local slot keys are unique per USER function; loop-carried locals keep ONE. +_pyir_fn_id_stack: "list[Any]" = [] + + +def pyir_enter_fn(fn_id: Any) -> None: + """Push *fn_id* as the current user-function scope and open the place-ledger + 'fn' frame: one scope_id per instrumented user-function activation.""" + _pyir_fn_id_stack.append(fn_id) + scope_id = next(_SCOPE_ID_COUNTER) + _PYIR_SCOPE_STACK.append(_ScopeFrame(scope_id, "fn")) + # F-SPEC: the first 'fn' scope after trace intake is the entry body; its + # parameter locals adopt the staged argument roots. + if _PYIR_SPEC_PENDING_PARAM_ROOTS: + for name, root in _PYIR_SPEC_PENDING_PARAM_ROOTS: + _PYIR_SPEC_LOCAL_ROOTS[("local", scope_id, name)] = root + _PYIR_SPEC_PENDING_PARAM_ROOTS.clear() + + +def pyir_exit_fn() -> None: + """Pop the current user-function scope (in the function body's ``finally``).""" + if _pyir_fn_id_stack: + _pyir_fn_id_stack.pop() + if _PYIR_SCOPE_STACK: + _PYIR_SCOPE_STACK.pop() + + +class _Skip: + """Sentinel type: the subscript-write container is not a tracked dict, so + the AST-injected pre/post hooks skip instrumentation. Compare by identity + against the single ``_PYIR_SKIP`` instance below.""" + + __slots__ = () + + def __repr__(self) -> str: + return "_PYIR_SKIP" + + +_PYIR_SKIP: "_Skip" = _Skip() + +# The extraction walk's DECLARED owner objects (innermost last): an unpaired +# read under the walk resolves its place row against these holders only. +_EXTRACTION_WALK_OWNERS: "list[Any]" = [] + + +_PYIR_READ_SIMPLE_SENTINEL = _Sentinel("read-simple miss") + + +# Superseded-generation detection: one live cell per leaf place survives a +# whole-object rebind, so stamped non-current generations must refuse reads. + +# id(non-current generation) -> stamp record ({"target", "site", "cells", +# "wrapper_refs"}); leaf-cell store versions recorded before the rebind lands. +_SUPERSEDED_GENERATIONS: "dict[int, dict]" = {} +# id(leaf wrapper) -> (stamp record, slot_name) for the superseded generation's +# own staged leaves at stamp time; catches owner-less auto-load reads. +_SUPERSEDED_LEAF_WRAPPERS: "dict[int, tuple]" = {} +# Pin for the stamp's id key: weakref whose callback drops the stamp, or a +# strong ref for non-weakrefable objects so the id cannot be recycled. +_GEN_OBJ_KEEPALIVE: "dict[int, Any]" = {} +_PYIR_TRACE_KEEPALIVES.append(_GEN_OBJ_KEEPALIVE) +# Binding-place slot key -> rebind-generation counter, bumped at every +# whole-object rebind of that place (m2m and Python-phase polarity alike). +_BINDING_GEN: "dict[Any, int]" = {} +# Per bumped key: pre-rebind leaf-cell store-version snapshot + rebind site -- +# the baseline the alias-capture fire rule compares against. +_BINDING_REBIND_CELLS: "dict[Any, dict]" = {} +# Plain-local slot key -> {"root", "gen"}: a compound rooted at another +# binding place got bound to a bare local (the alias-capture read channel). +_ALIAS_CAPTURES: "dict[Any, dict]" = {} +# Bare local NAMES with live alias captures, so the read choke can pre-screen +# a dotted read's access root without a frame walk. +_ALIAS_CAPTURE_ROOT_NAMES: "set[str]" = set() +# id(compound) -> the binding-place slot it was FIRST bound at (its root); +# feeds alias capture and the m2m generation bump. +_COMPOUND_BINDING_SLOTS: "dict[int, Any]" = {} +# Reentrancy depth: walk-internal per-field assigns and framework carry run +# with generation events and checks suppressed. +_PYIR_GEN_EVENT_SUPPRESS: "list[int]" = [0] + +# Region-conditional attr first-defs: (id(owner), slot_name) -> {"arms", "site"}; +# a read whose arm path does not extend the first-def's raises. Cleared at exit. +_PYIR_CF_ATTR_FIRST_DEFS: "dict[tuple, dict]" = {} -_MODULE_OPS = frozenset(("builtin.module", "gpu.module")) -__all__ = [name for name in list(globals()) if not name.startswith("__")] +# Static export surface (regenerate with scripts/gen_pyir_all.py): the +# module's OWN names -- its namespace minus what lower chain modules +# already export and minus the generator's INTERNAL_ONLY registry. +__all__ = [ + "inspect", + "sys", + "types", + "_weakref", + "TYPE_CHECKING", + "Any", + "Callable", + "_Literal", + "_NamedTuple", + "Optional", + "is_inside_staged_cf", + "is_inside_locally_staged_cf", + "is_inside_constexpr_loop", + "constexpr_scope_under_staged_cf", + "current_staged_cf_depth", + "_is_staged_value", + "assign_meta_staged_check", + "DSLRuntimeError", + "DSLUserCodeError", + "is_auto_m2s_enabled", + "is_pyir_enabled", + "DiagId", + "ir", + "log", + "pyir", + "ub", + "_PYIR_MODE_BY_PREFIX", + "_pyir_register_mode_fact", + "_pyir_mixed_mode_hint", + "_SCF_REGION_NAMES", + "_NON_DOT_FUNC_ENTRY_OPS", + "_MODULE_OPS", + "_meta_uses", + "_meta_idempotent_write_anchors", + "_slot_refs", + "_slot_templates", + "_PYIR_LAST_NOTIN_COMPARE", + "_PYIR_FOLD_FIRSTDEF_STACK", + "CF_ATTR_FIRST_DEF", + "_PYIR_GATE_ONCE_INIT_SLOTS", + "_PYIR_BOUNDARY_CALLEE_DEPTH", + "_slot_first_def_inside_cf", + "_slot_first_def_block", + "_slot_first_def_block_any", + "_slot_first_def_depth", + "_slot_first_def_depth_any", + "_slot_binding_depth", + "_pyir_open_loop_body_blocks", + "_pyir_bump_region_epoch", + "_pyir_current_region_epoch", + "_pyir_region_entry_push", + "_pyir_region_entry_pop", + "_pyir_region_entry_clock", + "_pyir_region_entry_watermark", + "_SLOT_REGISTRY", + "_OWNER_KEEPALIVE", + "_PYIR_PROMOTED_PLACE_LEAVES", + "_Sentinel", + "_NO_CONST_VALUE", + "_WATCHED_CONTAINER_ADOPTIONS", + "_WATCHED_CONTAINER_KEEPALIVE", + "_PYIR_SPEC_CONTAINER_CHAIN", + "_WATCHED_CONTAINER_WALKED", + "_WATCHED_CONTAINER_WALKED_KEEPALIVE", + "_WATCHED_CONTAINER_BYPASS", + "_WATCHED_DICT_READ_HOOK", + "_WATCHED_DICT_WRITE_HOOK", + "_WATCHED_DICT_GET_MISS_HOOK", + "_WATCHED_DICT_MUTATOR_HOOK", + "_WATCHED_DICT_ITER_HOOK", + "_WATCHED_LIST_READ_HOOK", + "_WATCHED_LIST_WRITE_HOOK", + "_WATCHED_LIST_MUTATOR_HOOK", + "_WATCHED_CONTAINER_HELD_OBJECTS", + "_WATCHED_CONTAINER_HELD_KEEPALIVE", + "_PYIR_TRACE_KEEPALIVES", + "_SLOT_STORE_ATTR", + "_PYIR_SLOT_HOLDERS", + "_PYIR_CANDIDATE_HOLDERS", + "_PYIR_SIGHTING_SETTLED", + "_PYIR_CANDIDATE_REGISTRY_ACTIVE", + "_PYIR_WRITE_CLOCK", + "_PYIR_HOLDER_WRITE_STAMPS", + "_PYIR_GATHER_SEGMENTS", + "_pyir_note_holder_write", + "_pyir_setattr_raw", + "_pyir_delattr_raw", + "_pyir_declare_write_funnel", + "_pyir_wrapper_write_funneled", + "_pyir_getattr_reads_storage", + "_pyir_plain_storage_setattr", + "_PYIR_SELF_LEAF_ATTR", + "_NUMERIC_LEAF_NONE", + "_NUMERIC_LEAF_SSA", + "_NUMERIC_LEAF_META", + "_PYIR_LOOP_ITER_ARGS_ATTR", + "_PYIR_TUPLE_LEAF_ATTR", + "_PYIR_LIST_MUTATORS", + "_PYIR_DICT_MUTATORS", + "_PYIR_SET_MUTATORS", + "_PYIR_DEQUE_MUTATORS", + "_PYIR_LIST_INSERT_MUTATORS", + "_PYIR_DICT_INSERT_MUTATORS", + "_SCOPE_ID_COUNTER", + "_TOKEN_COUNTER", + "_ScopeFrame", + "_PYIR_SCOPE_STACK", + "_PlaceSeg", + "_place_seg_child", + "IdentityKeyedWeakTable", + "_OWNER_TOKENS", + "_PYIR_REBUILD_PROTOCOL_DEPTH", + "_PYIR_TOKEN_BORN_CLASS", + "_PLACE_REGISTRY", + "_PYIR_SUPERSEDED_PLACE_ROWS", + "_OWNER_PLACE_PREFIX", + "_ROOT_NAME_TOKENS", + "_PYIR_SCOPE_CELL_TOKENS", + "_PYIR_CELL_HOME_BINDING", + "_PYIR_SCOPE_NONLOCAL_NAMES", + "_CE_BINDING_OWNER", + "_EMISSION_SELF_CHECK_LOG", + "_EMISSION_SELF_CHECK_LOG_CAP", + "_PYIR_SPEC_TOKEN_ROOTS", + "_PYIR_SPEC_LOCAL_ROOTS", + "_PYIR_SPEC_PENDING_PARAM_ROOTS", + "_PYIR_SPEC_UNROOTED_TOKENS", + "_PYIR_SPEC_TRACE_BORN_TOKENS", + "_PYIR_SPEC_TRACE_BORN_PINS", + "_PYIR_SPEC_RECORD", + "_PYIR_SPEC_COMPLETE", + "_PYIR_SPEC_ENTRY_FUNC", + "_PYIR_SPEC_ENTRY_SIG", + "_PYIR_SPEC_ENTRY_RECEIVER", + "_SPEC_RECEIVER_DEAD", + "_SpecTraceExitObject", + "_SpecContainerSnapshot", + "_SPEC_ATTR_ABSENT", + "_SPEC_ATTR_PRESENT", + "_PYIR_SPEC_SEALED", + "_reset_spec_state", + "_reset_scope_state", + "pyir_enter_fn", + "pyir_exit_fn", + "_Skip", + "_PYIR_SKIP", + "_EXTRACTION_WALK_OWNERS", + "_PYIR_READ_SIMPLE_SENTINEL", + "_SUPERSEDED_GENERATIONS", + "_SUPERSEDED_LEAF_WRAPPERS", + "_GEN_OBJ_KEEPALIVE", + "_BINDING_GEN", + "_BINDING_REBIND_CELLS", + "_ALIAS_CAPTURES", + "_ALIAS_CAPTURE_ROOT_NAMES", + "_COMPOUND_BINDING_SLOTS", + "_PYIR_GEN_EVENT_SUPPRESS", + "_PYIR_CF_ATTR_FIRST_DEFS", +] diff --git a/python/CuTeDSL/cutlass/base_dsl/pyir_threading.py b/python/CuTeDSL/cutlass/base_dsl/pyir_threading.py deleted file mode 100644 index 3c5cfbf56a..0000000000 --- a/python/CuTeDSL/cutlass/base_dsl/pyir_threading.py +++ /dev/null @@ -1,186 +0,0 @@ -# SPDX-FileCopyrightText: Copyright (c) 2025 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# SPDX-License-Identifier: LicenseRef-NvidiaProprietary -# -# Use of this software is governed by the terms and conditions of the -# NVIDIA End User License Agreement (EULA), available at: -# https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/license.html -# -# Any use, reproduction, disclosure, or distribution of this software -# and related documentation outside the scope permitted by the EULA -# is strictly prohibited. - - -"""PyIR runtime -- threading layer; see facade for the public surface.""" - -from .pyir_corewalk import * # noqa: F401,F403 (re-export lower layers up the chain) - -# -- BEGIN explicit imports for the type checker (do not edit the list by hand; -# it mirrors names the chain re-exports at runtime via the wildcard + dynamic -# ``__all__`` above, which a static type checker cannot evaluate -- so every -# name is also imported explicitly from the layer that DEFINES it). Purely -# additive: the wildcard import stays the runtime source of truth. -from .pyir_state import ( # noqa: F401 - Any, - _meta_uses, - _slot_first_def_block, - _slot_first_def_inside_cf, - _slot_refs, - ir, - log, - pyir, -) -from .pyir_core import ( # noqa: F401 - _ancestor_op_in_block, - _emit_constant_at_ip, - _get_function_entry_block, - _raise_on_fold_witness, - _replace_value_uses, -) -# -- END explicit imports for the type checker - - -def _meta_promote_slot( - slot_key: Any, - initial_py_value: Any, - target_name: str | None = None, - filename: str | None = None, - lineno: int | None = None, -) -> "ir.Value | None": - """Promote a slot to D1 tracking. - - Creates a ``pyir.ref`` at the enclosing function's entry block - initialized with *initial_py_value*, then walks every previously - baked ``arith.constant`` recorded under *slot_key* and replaces its - uses with a freshly-emitted ``pyir.load %ref``. The replaced - constants become dead and DCE cleans them up; downstream arith ops - pick up the load via SSA edges automatically -- no per-derivation - rewriting needed. - - Emits a ``DSLWarning`` (WarnId.PHASE_AUTO_PROMOTED_TO_STAGED) via - ``report_warning`` so users can see which variable was promoted, its - original value, and the file:line of the offending mutation. Shares the - catalog entry with the Mp→S warning in ``pyir_read`` so both promotion - paths surface identically. - - Returns the new ``pyir.ref`` SSA value, or ``None`` if there is no - enclosing function (e.g. tracing happens outside an MLIR context). - """ - if pyir is None: - return None - # P6: a slot whose constexpr fold already decided CPython control flow - # cannot be promoted -- its reads would become runtime loads while the - # baked branch stays hard-wired (silent miscompile). Refuse here, at the - # single promotion entry. - _raise_on_fold_witness(slot_key, initial_py_value, target_name, filename, lineno) - from .pyir_runtime import _get_function_entry_block, _ancestor_op_in_block - - entry_block = _get_function_entry_block() - if entry_block is None: - return None - - with ir.InsertionPoint.at_block_begin(entry_block): - initial_ir = _emit_constant_at_ip(initial_py_value) - ref = pyir.ref(initial_ir) - _slot_refs[slot_key] = ref - - # User-visible warning: a Python (M) value just became staged-tracked - # (ref + load/store). Show target, original value, file:line so - # users can locate and (optionally) hoist the init out of CF. - if target_name is not None: - from .diagnostics import WarnId, report_warning - - if isinstance(initial_py_value, bool): - promoted_type = "Boolean" - elif isinstance(initial_py_value, int): - promoted_type = "Int32" - elif isinstance(initial_py_value, float): - promoted_type = "Float32" - else: - promoted_type = type(initial_py_value).__name__ - # Warn about the promotion AND the stale-local risk: the slot is now - # tracked, but the caller's Python binding for ``target_name`` still - # holds the original value, so a read OUTSIDE this region could see the - # stale value -- the catalog message explains this to the author. - report_warning( - WarnId.PHASE_AUTO_PROMOTED_TO_STAGED, - filename=filename, - lineno=lineno, - stacklevel=4, - var=target_name, - value=repr(initial_py_value), - type=promoted_type, - ) - - # Per-iteration reset for the FIRST unrolled iteration's first-def. - # Only fires when the user's first-def was INSIDE staged CF (e.g. - # ``acc = 0.0`` at the top of an unrolled constexpr loop body whose - # outer parent is a staged ``scf.for``/``scf.while``). In that - # case the function-entry ref init carries the primitive across - # iter 0 of the enclosing runtime loop for free, but on iter 1+ the - # ref becomes an ``iter_arg`` carrying the yielded value from iter 0 - # and the user's reset at the top of the body is lost (the wrap-as- - # _WatchedM first-def emits no IR). Insert ``pyir.store(initial_const, - # %ref)`` BEFORE the earliest baked use so every runtime iteration - # observes the reset. The follow-on ``replaceAllUsesWith`` loop - # converts that bake into ``pyir.load %ref`` -- our store dominates - # the load via IR order. - # - # NOT fired when the first-def is OUTSIDE staged CF (e.g. - # ``self.val = True`` in a plain ``__init__``), because the user - # then intends the slot to persist across iterations of any runtime - # loop and a per-iteration reset would clobber subsequent mutations - # back to the initial literal. - if _slot_first_def_inside_cf.get(slot_key, False): - uses = _meta_uses.get(slot_key, []) - if uses: - first_use = uses[0] - try: - owner = first_use.owner - if not isinstance(owner, ir.Block): - defining_op = ( - owner - if isinstance(owner, ir.Operation) - else getattr(owner, "operation", owner) - ) - # L103: the earliest baked use may live inside a NESTED - # runtime loop entered after the first-def. Hoist the - # reset to the first-def's block so it runs once per - # first-def-level iteration (before the nested loop), - # not on every inner iteration. - anchor_op = _ancestor_op_in_block( - defining_op, _slot_first_def_block.get(slot_key) - ) - with ir.InsertionPoint(anchor_op): - reset_ir = _emit_constant_at_ip(initial_py_value) - pyir.store(reset_ir, ref) - except Exception as exc: - log().info( - "[_meta_promote_slot] %s: reset-store insert failed: %s", - slot_key, - exc, - ) - - for const_val in _meta_uses.pop(slot_key, []): - try: - owner = const_val.owner - if isinstance(owner, ir.Block): - continue # block argument has no "before" insertion point - defining_op = ( - owner - if isinstance(owner, ir.Operation) - else getattr(owner, "operation", owner) - ) - with ir.InsertionPoint(defining_op): - loaded = pyir.load(ref) - if not _replace_value_uses(const_val, loaded): - log().info( - "[_meta_promote_slot] %s: replace_all_uses_with missing on ir.Value", - slot_key, - ) - except Exception as exc: - log().info("[_meta_promote_slot] %s: replace failed: %s", slot_key, exc) - - return ref - - -__all__ = [name for name in list(globals()) if not name.startswith("__")] diff --git a/python/CuTeDSL/cutlass/base_dsl/runtime/cuda.py b/python/CuTeDSL/cutlass/base_dsl/runtime/cuda.py index 11c6afd57d..386fcc5626 100644 --- a/python/CuTeDSL/cutlass/base_dsl/runtime/cuda.py +++ b/python/CuTeDSL/cutlass/base_dsl/runtime/cuda.py @@ -17,7 +17,7 @@ from dataclasses import dataclass from typing import Any from enum import IntEnum -import numpy as np +import ctypes import os import cuda.bindings.driver as cuda @@ -434,11 +434,11 @@ def initialize_cuda_context(device_id: int = 0, flags: int = 0) -> Any: driver_version = get_driver_version() - # Check the CUDA driver version works for the installed cuda-python package + # Check the CUDA driver version works for the installed cuda-bindings package if driver_version < 13000 and cuda.CUDA_VERSION >= 13000: raise DSLRuntimeError( - f"CUDA driver version {driver_version} is below the minimum required version for the installed cuda-python package {cuda.CUDA_VERSION}.", - suggestion=f"Consider updating your NVIDIA driver to version 580 or above. Or install cuda-python package with version 12.9 or below.", + f"CUDA driver version {driver_version} is below the minimum required version for the installed cuda-bindings package {cuda.CUDA_VERSION}.", + suggestion=f"Consider updating your NVIDIA driver to version 580 or above. Or install cuda-bindings package with version 12.9 or below.", ) # Check if a valid CUDA context already exists (e.g., created by PyTorch or @@ -542,11 +542,13 @@ def load_cubin_module_data(cubin_data: bytes) -> Any: :rtype: cuda.CUmodule :raise DSLRuntimeError: If the CUDA operation fails. """ - # Load module data - _log().info(f"cuModuleLoadData {np.char.array(cubin_data).ctypes.data}") - module = checkCudaErrors( - cuda.cuModuleLoadData(np.char.array(cubin_data).ctypes.data) - ) + # Load module data. Keep the buffer alive in a local across the + # cuModuleLoadData call: `address` is only a raw address into it, so a bare + # temporary could be garbage-collected before the driver reads it. + _cubin_buf = ctypes.create_string_buffer(cubin_data) + address = ctypes.addressof(_cubin_buf) + _log().info(f"cuModuleLoadData {address}") + module = checkCudaErrors(cuda.cuModuleLoadData(address)) return module @@ -607,14 +609,14 @@ def load_library_data(cubin_data: bytes | int) -> Any: :raise DSLRuntimeError: If the CUDA operation fails. """ # Load module data - # Keep the numpy array alive in a local (`_cubin_arr`) across the - # cuLibraryLoadData call: `.ctypes.data` is only a raw address into the - # array's buffer, so if the array were a bare temporary it could be + # Keep the buffer alive in a local (`_cubin_buf`) across the + # cuLibraryLoadData call: the address is only a raw address into the + # buffer, so if the buffer were a bare temporary it could be # garbage-collected before the driver reads the address (use-after-free). - _cubin_arr = None + _cubin_buf = None if isinstance(cubin_data, bytes): - _cubin_arr = np.char.array(cubin_data) - cubin_data = _cubin_arr.ctypes.data + _cubin_buf = ctypes.create_string_buffer(cubin_data) + cubin_data = ctypes.addressof(_cubin_buf) _log().info(f"cuLibraryLoadData {cubin_data!r}") library = checkCudaErrors( diff --git a/python/CuTeDSL/cutlass/base_dsl/runtime/jit_arg_adapters.py b/python/CuTeDSL/cutlass/base_dsl/runtime/jit_arg_adapters.py index 50a0c087a5..58b566edfe 100644 --- a/python/CuTeDSL/cutlass/base_dsl/runtime/jit_arg_adapters.py +++ b/python/CuTeDSL/cutlass/base_dsl/runtime/jit_arg_adapters.py @@ -138,8 +138,9 @@ class JitArgAdapterRegistry: # Adapters keyed by fully-qualified type name ("module.QualName") for # types whose defining module is too expensive to import at registration - # time (e.g. torch). Promoted into jit_arg_adapter_registry on first - # lookup of an instance, which can only exist once the module is loaded. + # time (e.g. torch). Cached in jit_arg_adapter_registry for each concrete + # type on first lookup of an instance, which can only exist once the module + # is loaded. lazy_jit_arg_adapter_registry: dict[str, Any] = {} # Top-level module names of the lazy registrations, so lookups of @@ -179,7 +180,9 @@ class MyAdapterForMyPythonType: "module.QualName" string instead, so registration never imports the defining module (e.g. torch). An instance of the type can only reach a JIT function after the application has imported its module, so the - adapter is promoted to the concrete-type registry on first lookup. + adapter is cached for the concrete type on first lookup. Named base + classes also match their subclasses, allowing registration against a + stable public type when implementations use private concrete types. """ if python_type is None: @@ -240,11 +243,15 @@ def decorator(adapter: Any) -> Any: @classmethod def _promote_lazy_adapter(cls, python_type: type) -> Any: - type_qualname = f"{python_type.__module__}.{python_type.__qualname__}" - adapter = cls.lazy_jit_arg_adapter_registry.pop(type_qualname, None) - if adapter is not None: - cls.jit_arg_adapter_registry[python_type] = adapter - return adapter + for candidate_type in python_type.__mro__: + type_qualname = ( + f"{candidate_type.__module__}.{candidate_type.__qualname__}" + ) + adapter = cls.lazy_jit_arg_adapter_registry.get(type_qualname) + if adapter is not None: + cls.jit_arg_adapter_registry[python_type] = adapter + return adapter + return None @classmethod @contextmanager @@ -285,6 +292,19 @@ def get_registered_adapter(cls, arg: object) -> Any: exists, but reports ambiguity instead of silently choosing an adapter based on module import order. """ + # Transparency protocol (F-TRANSPARENT): a tracing observation wrapper + # that declares a plain counterpart resolves and runs its adapter AS + # that plain value. The wrapper type is a host-trace artifact and is + # never registered, so an exact-type lookup would miss the plain type's + # adapter and the argument would reach the JIT boundary unconverted. + # Same consult as ``utils.tree_utils._tree_flatten``. + if getattr(type(arg), "__pyir_plain_type__", None) is not None: + plain_view: Any = arg.__pyir_plain_view__() # type: ignore[attr-defined] + plain_adapter = cls.get_registered_adapter(plain_view) + if plain_adapter is None: + return None + return lambda wrapper: plain_adapter(wrapper.__pyir_plain_view__()) + python_type = type(arg) resolved_scope = cls._active_scope.get() adapter = None @@ -299,8 +319,10 @@ def get_registered_adapter(cls, arg: object) -> Any: if ( adapter is None and cls.lazy_jit_arg_adapter_registry - and python_type.__module__.partition(".")[0] - in cls._lazy_adapter_module_roots + and not cls._lazy_adapter_module_roots.isdisjoint( + candidate_type.__module__.partition(".")[0] + for candidate_type in python_type.__mro__ + ) ): adapter = cls._promote_lazy_adapter(python_type) diff --git a/python/CuTeDSL/cutlass/base_dsl/tvm_ffi_builder/mlir_builder.py b/python/CuTeDSL/cutlass/base_dsl/tvm_ffi_builder/mlir_builder.py index 76f39c6e30..3f03d2b9ea 100644 --- a/python/CuTeDSL/cutlass/base_dsl/tvm_ffi_builder/mlir_builder.py +++ b/python/CuTeDSL/cutlass/base_dsl/tvm_ffi_builder/mlir_builder.py @@ -114,6 +114,8 @@ class MLIRBuilder(MLIRTypeBuilder): but not set the insersion point. """ + MLIR_DYNAMIC_INDEX = -(2**31) + def __init__(self) -> None: """Initialize the MLIR builder.""" super().__init__() @@ -265,7 +267,11 @@ def address_of(self, name: str, tp: ir.Type) -> ir.Value: return llvm.AddressOfOp(tp, name).result def getelementptr( - self, ptr: ir.Value, constant_indices: Sequence[int], elem_type: ir.Type + self, + ptr: ir.Value, + constant_indices: Sequence[int], + elem_type: ir.Type, + dynamic_indices: Sequence[ir.Value] = [], ) -> ir.Value: """Create a getelementptr operation. @@ -273,10 +279,13 @@ def getelementptr( ---------- ptr : ir.Value The pointer to the element. - indices : Sequence[ir.Value] - The indices to the element. + constant_indices : Sequence[int] + The indices to the element. Use ``MLIR_DYNAMIC_INDEX`` as a + placeholder for an index taken from ``dynamic_indices`` instead. elem_type : ir.Type The type of the element. + dynamic_indices : Sequence[ir.Value], optional + Runtime index values. Returns ------- @@ -287,7 +296,7 @@ def getelementptr( return llvm.getelementptr( self.ptr_type, ptr, - [], + dynamic_indices, raw_constant_indices=ir.DenseI32ArrayAttr.get(constant_indices), elem_type=elem_type, **self.get_element_extra_kwargs, @@ -298,12 +307,21 @@ def getelementptr( return llvm.getelementptr( self.ptr_type, ptr, - [], + dynamic_indices, raw_constant_indices=ir.DenseI32ArrayAttr.get(constant_indices), elem_type=elem_type, **self.get_element_extra_kwargs, ) + def offset_ptr_bytes(self, ptr: ir.Value, byte_offset: ir.Value) -> ir.Value: + """Displace ``ptr`` by a runtime number of bytes.""" + return self.getelementptr( + ptr, + [self.MLIR_DYNAMIC_INDEX], + self.i8_type, + dynamic_indices=[byte_offset], + ) + def return_(self, ret: Optional[ir.Value] = None) -> None: """Create a return statement.""" llvm.return_(arg=ret) diff --git a/python/CuTeDSL/cutlass/base_dsl/tvm_ffi_builder/tvm_ffi_builder.py b/python/CuTeDSL/cutlass/base_dsl/tvm_ffi_builder/tvm_ffi_builder.py index 6d0d4f802b..7e103bb8fc 100644 --- a/python/CuTeDSL/cutlass/base_dsl/tvm_ffi_builder/tvm_ffi_builder.py +++ b/python/CuTeDSL/cutlass/base_dsl/tvm_ffi_builder/tvm_ffi_builder.py @@ -2079,6 +2079,7 @@ def decode_param_tensor( device_id = self.load_dltensor_device_id(dl_tensor_ptr) ndim = self.load_dltensor_ndim(dl_tensor_ptr) byte_offset = self.load_dltensor_byte_offset(dl_tensor_ptr) + data = self.offset_ptr_bytes(data, byte_offset) # check data alignment if specified if param.data_alignment is not None: @@ -2176,16 +2177,6 @@ def dtype_equal() -> ir.Value: "Mismatched Tensor ", arg_context, f", expected dtype={param.dtype}" ), ) - # check byte_offset - # Break error message into reusable parts for better string deduplication - current_block = self.check_condition( - current_block, - lambda: self.equal(byte_offset, self.i64(0)), - "ValueError", - self._arg_err( - "Mismatched Tensor ", arg_context, ", expected byte_offset=0" - ), - ) with ir.InsertionPoint(current_block): shape = self.load_dltensor_shape(dl_tensor_ptr) diff --git a/python/CuTeDSL/cutlass/base_dsl/typing.py b/python/CuTeDSL/cutlass/base_dsl/typing.py index 99b64fb2fe..c8965124bd 100644 --- a/python/CuTeDSL/cutlass/base_dsl/typing.py +++ b/python/CuTeDSL/cutlass/base_dsl/typing.py @@ -11,9 +11,9 @@ import ctypes import math +import struct from abc import abstractmethod from itertools import chain -import numpy as np import operator from typing import ( TYPE_CHECKING, @@ -28,6 +28,7 @@ Union, Any, cast as tcast, + SupportsIndex, Type, TypeVar, overload, @@ -48,7 +49,18 @@ from . import AddressSpace -from .pyir_runtime import _WatchedM +from .pyir_runtime import ( + _PYIR_CANDIDATE_REGISTRY_ACTIVE, + _PYIR_SCOPE_STACK, + _WatchedM, + _pyir_declare_write_funnel, + _pyir_note_holder_write, + _pyir_record_external_payload_consumption, + _pyir_record_staged_identity_hashed, + _pyir_refresh_cell_read, + _pyir_register_candidate_holder, + _pyir_setattr_raw, +) # ============================================================================= # Dynamic Expression Protocol @@ -370,16 +382,17 @@ class NumericMeta(DslType): For unpacked dtypes this is the element width. Packed view dtypes use the width of one packed tensor element. :type width: int - :param np_dtype: Corresponding NumPy dtype - :type np_dtype: numpy.dtype, optional + :param np_dtype_name: Name of the corresponding NumPy scalar type, or None + when NumPy has no matching type + :type np_dtype_name: str, optional :param mlir_type: Corresponding MLIR type :type mlir_type: Any, optional :param is_abstract: Whether the type is abstract, defaults to False :type is_abstract: bool, optional :ivar width: Bit width of the numeric type :type width: int - :ivar _np_dtype: Corresponding NumPy dtype - :type _np_dtype: Union[numpy.dtype, None] + :ivar _np_dtype_name: Name of the corresponding NumPy scalar type + :type _np_dtype_name: Union[str, None] :property numpy_dtype: Returns the corresponding NumPy dtype :rtype numpy_dtype: numpy.dtype @@ -390,7 +403,7 @@ class NumericMeta(DslType): # Placeholder type _mlir_type = Any - _np_dtype: Optional[type] + _np_dtype_name: Optional[str] def __new__( cls, @@ -398,7 +411,7 @@ def __new__( bases: tuple, attrs: dict, width: int = 8, - np_dtype: Optional[type] = None, + np_dtype_name: Optional[str] = None, mlir_type: Optional[Callable[[], ir.Type]] = None, is_abstract: bool = False, **kwargs: Any, @@ -428,7 +441,7 @@ def _new_from_mlir_values(self: "Numeric", values: list[ir.Value]) -> "Numeric": new_cls.width = width new_cls.bytes = max(1, (width + 7) // 8) - new_cls._np_dtype = np_dtype + new_cls._np_dtype_name = np_dtype_name return new_cls def n_bytes(cls, n_elements: int) -> int: @@ -438,7 +451,20 @@ def n_bytes(cls, n_elements: int) -> int: @property def numpy_dtype(cls) -> Optional[type]: - return cls._np_dtype + """Return the NumPy scalar type for this dtype, or None if it has none. + + NumPy is an optional dependency, so it is imported on first access + rather than at module load: nothing else in the type system needs it. + + :raises ModuleNotFoundError: If NumPy is not installed and this dtype + does have a NumPy counterpart. + """ + if cls._np_dtype_name is None: + return None + + import numpy + + return getattr(numpy, cls._np_dtype_name, None) @property @abstractmethod @@ -534,6 +560,20 @@ def cast( return res +_INTEGER_DTYPE_NAMES: dict[tuple[int, bool], str] = { + (8, True): "int8", + (16, True): "int16", + (32, True): "int32", + (64, True): "int64", + (8, False): "uint8", + (16, False): "uint16", + (32, False): "uint32", + (64, False): "uint64", +} + +_FLOAT_DTYPE_NAMES = frozenset({"float16", "float32", "float64"}) + + # Option 1: use ir.Value as base # class IntegerMeta(DslType, type(ir.Value)): class IntegerMeta(NumericMeta): @@ -553,9 +593,12 @@ class IntegerMeta(NumericMeta): """ signed: bool - # Value range that ``_np_dtype`` stores exactly, or None when there is no - # numpy dtype to store into. See ``Integer.__init__``. + # Value range this type stores exactly, or None when the width has no + # matching C integer to cast to. See ``Integer.__init__``. _exact_range: Optional[tuple[int, int]] + # (width, signedness) of the C integer this type casts to, or None when + # there is none. See ``_float_to_int``. + _cast_key: Optional[tuple[int, bool]] def __new__( cls, @@ -567,16 +610,9 @@ def __new__( mlir_type: Optional[Callable[[], ir.Type]] = None, is_abstract: bool = False, ) -> Any: - if width == 1: - np_dtype = np.bool_ - elif width == 128: - np_dtype = None - elif width == 4: - np_dtype = None - elif signed: - np_dtype = getattr(np, f"int{width}") - else: - np_dtype = getattr(np, f"uint{width}") + np_dtype_name = ( + "bool_" if width == 1 else _INTEGER_DTYPE_NAMES.get((width, signed)) + ) def _c_pointers(self: "Integer") -> list[ctypes.c_void_p]: if width == 1: @@ -592,22 +628,41 @@ def _c_pointers(self: "Integer") -> list[ctypes.c_void_p]: "__c_pointers__": _c_pointers, } new_cls = super().__new__( - cls, name, bases, attrs | new_attrs, width, np_dtype, mlir_type, is_abstract + cls, + name, + bases, + attrs | new_attrs, + width, + np_dtype_name, + mlir_type, + is_abstract, ) new_cls.signed = signed # Precomputed once per type so ``Integer.__init__`` can range-check without - # rebuilding the bounds on every construction. These track ``np_dtype`` - # rather than ``width``, because the numpy cast is what they must agree - # with, and a few subclasses patch ``width`` afterwards. - if np_dtype is None: + # rebuilding the bounds on every construction. These track the dtype + # rather than the ``width`` attribute, because the dtype's C cast is what + # they must agree with, and a few subclasses patch ``width`` afterwards -- + # hence the local ``width``/``signed`` parameters below rather than + # ``cls.width``. The final branches are reached only for widths 8/16/32/64, + # where the dtype is exactly ``int{width}``/``uint{width}``, so + # two's-complement arithmetic on the parameters agrees with it by + # construction. + if np_dtype_name is None: new_cls._exact_range = None elif width == 1: - # np.bool_ folds everything nonzero to True, so only 0 and 1 survive + # bool folds everything nonzero to True, so only 0 and 1 survive # the cast unchanged. new_cls._exact_range = (0, 1) + elif signed: + new_cls._exact_range = (-(2 ** (width - 1)), 2 ** (width - 1) - 1) else: - info = np.iinfo(np_dtype) # type: ignore[type-var] - new_cls._exact_range = (int(info.min), int(info.max)) + new_cls._exact_range = (0, 2**width - 1) + # Boolean is declared signed with width 1 yet folds to (0, 1), and it + # has no C integer to cast to, so it is excluded alongside the widths + # that have no dtype at all. + new_cls._cast_key = ( + None if (np_dtype_name is None or width == 1) else (width, signed) + ) return new_cls def __str__(cls) -> str: @@ -683,14 +738,15 @@ def __new__( mlir_type: Optional[Callable[[], ir.Type]] = None, is_abstract: bool = False, ) -> Any: - np_dtype = getattr(np, name.lower(), None) + lowered_name = name.lower() + np_dtype_name = lowered_name if lowered_name in _FLOAT_DTYPE_NAMES else None new_cls = super().__new__( cls, name, bases, attrs, width, - np_dtype, + np_dtype_name, mlir_type, is_abstract, ) @@ -915,6 +971,59 @@ def _promote_integer( return a.to(b.dtype), b, b.dtype +def _pyir_mark_fold_srcs(result: Any, operands: "tuple[Any, ...]") -> None: + """Stamp on *result* the payload-backed Numeric operands (plus their own + recorded sources) whose trace-time values decided it; consumed by the + staged-literal predicate-fold witness (see pyir_core). + + Gated on an open PyIR trace scope (the same check the consumption + recorder uses): outside a trace there is no witness to consume the + stamp, so scalar folds stay untouched.""" + if not _PYIR_SCOPE_STACK: + return + try: + srcs: list = [] + for o in operands: + if isinstance(o, Numeric) and type(getattr(o, "value", None)) in ( + bool, + int, + float, + ): + if all(s is not o for s in srcs): + srcs.append(o) + for s in getattr(o, "_pyir_fold_srcs", ()): + if all(x is not s for x in srcs): + srcs.append(s) + if srcs: + _pyir_setattr_raw(result, "_pyir_fold_srcs", tuple(srcs)) + except Exception: + pass + + +def _pyir_numeric_protocol_gate(proto: str, *operands: Any) -> None: + """Refuse a numeric protocol that has no staged form when any operand is a + staged (runtime) value; payload operands proceed to the fold arm.""" + for o in operands: + staged = isinstance(o, ArithValue) or ( + isinstance(o, Numeric) + and not isinstance(o.__dict__.get("value"), (bool, int, float)) + ) + if staged: + raise DSLUserCodeError( + DiagId.PHASE_NUMERIC_PROTOCOL_ON_STAGED, + proto=proto, + what=type(o).__name__, + ) + + +def _pyir_numeric_payload(operand: Any) -> Any: + """The plain Python payload of a fold-arm operand (gate-checked Numeric + wrappers expose it as ``value``; plain primitives pass through).""" + if isinstance(operand, Numeric): + return operand.__dict__.get("value") + return operand + + def _binary_op( op: Callable[..., Any], promote_operand: bool = True, @@ -956,6 +1065,9 @@ def wrapper( ) -> Any: orig_lhs_type = type(lhs) orig_rhs_type = type(rhs) + # Captured BEFORE promotion rebinds lhs/rhs: the fold witness must key + # the ORIGINAL objects (promotion may mint twin wrappers). + orig_operands = (lhs, rhs) # When called directly with self and other ty = type(lhs) @@ -1009,7 +1121,11 @@ def wrapper( lhs_val, rhs_val = rhs_val, lhs_val res_val = op(lhs_val, rhs_val) - return res_type(res_val, loc=loc, ip=ip) + res = res_type(res_val, loc=loc, ip=ip) + # A Python-payload compute is a trace-time fold: propagate its sources. + if type(res_val) in (bool, int, float): + _pyir_mark_fold_srcs(res, orig_operands) + return res return wrapper @@ -1021,7 +1137,7 @@ class Numeric(metaclass=NumericMeta, is_abstract=True): implementing basic arithmetic operations. :param value: The value to store in the numeric type - :type value: Union[bool, int, float, Value] + :type value: Union[bool, int, float, Value, Numeric] :ivar value: The stored numeric value :vartype value: Union[bool, int, float, Value] @@ -1030,7 +1146,7 @@ class Numeric(metaclass=NumericMeta, is_abstract=True): # Injected by NumericMeta.__new__ on every concrete subclass. width: ClassVar[int] bytes: ClassVar[int] - _np_dtype: ClassVar[Optional[type]] + _np_dtype_name: ClassVar[Optional[str]] # TODO: Consider implementing MLIR style interface for PyIR mutable values. # Marker: MutableValue can track this type via pyir.ref/load/store. @@ -1038,14 +1154,47 @@ class Numeric(metaclass=NumericMeta, is_abstract=True): # Non-scalar DSL types (Array, Tensor) do NOT set this. _pyir_ref_supported = True + # The payload is a bool/int/float literal, an ``ir.Value`` or another + # Numeric depending on stage, and every reader narrows it itself. Declared + # here because the accessor that used to carry this annotation is bound + # onto the class only when PyIR is enabled. + value: Any + def __init__( self, - value: Union[bool, int, float, Value], + # Numeric is accepted: every concrete subclass constructor converts + # from Numeric inputs (e.g. Int2(Integer(0)) is a dataclass default), + # and type[Numeric] calls resolve against this signature. + value: Union[bool, int, float, Value, "Numeric"], *, loc: Optional[ir.Location] = None, ip: Optional[ir.InsertionPoint] = None, ) -> None: self.value = value + # PyIR value-wrapper MINT choke point: every scalar Numeric constructor + # chains through here, once per scalar kernel argument on the launch + # path included, so each guard below is spelled inline: a call whose + # only job is to return still costs a Python frame there. + + # A bare wrapper bound by an un-instrumented assignment never passes + # ``pyir_assign``/``pyir_read``; this sighting keeps it discoverable. + if _PYIR_CANDIDATE_REGISTRY_ACTIVE[0]: + _pyir_register_candidate_holder(self) + # Birth-context stamp: an SSA-backed wrapper belongs to the MLIR context + # that minted it (a live ``ir.Value`` keeps the context id stable). + + # The read funnels refuse a cross-context wrapper instead of + # dereferencing a handle whose defining op is gone. PyIR-only: the + # trace-close host-place restore keys this-compilation identity on the + # stamp, and only PyIR's read funnels consume it, so an unstamped + # wrapper in the baseline mode pays no per-mint guard work. The + # C-level type test leads so a literal payload never reaches the + # environment manager lookup. + if isinstance(value, ir.Value) and is_pyir_enabled(): + try: + self._pyir_birth_ctx = id(ir.Context.current) + except Exception: + pass def __str__(self) -> str: # Use member's pretty-str method if member object has method. @@ -1060,7 +1209,23 @@ def __repr__(self) -> str: return f"{self.__class__.__name__}({repr(self.value)})" def __hash__(self) -> int: - return hash(type(self).__class__) ^ hash(self.value) + v = self.__dict__.get("value") + if type(v) in (bool, int, float): + # Hashing the payload (dict key / set member) consumes it as structure. + _pyir_record_external_payload_consumption(self) + elif isinstance(v, ir.Value): + # Hashing a STAGED payload launders runtime identity into a plain + # Python hash; record it so the container key wall's write-site + # snapshot/compare can see the consumption (recording never raises). + _pyir_record_staged_identity_hashed() + return hash(type(self).__class__) ^ hash(v) + + def __reduce_ex__(self, protocol: SupportsIndex) -> Any: + v = self.__dict__.get("value") + if type(v) in (bool, int, float): + # Serialization consumes the payload as structure. + _pyir_record_external_payload_consumption(self) + return super().__reduce_ex__(protocol) @property def dtype(self) -> Type["Numeric"]: @@ -1169,7 +1334,30 @@ def to( if isinstance(self.value, (int, float, bool)): res = arith_helper.const(self.value, type(self), loc=loc, ip=ip) elif isinstance(self.value, ir.Value): - res = self.value + # Cross-compilation guard: a raw SSA minted under a DIFFERENT MLIR + # context belongs to a finalized compilation -- refuse loudly. + # An unstamped wrapper (baseline mode never stamps) skips the + # guard entirely, paying one attribute probe and no context id. + _birth: "int | None" = getattr(self, "_pyir_birth_ctx", None) + if _birth is not None: + _cur_ctx: "int | None" + try: + _cur_ctx = id(ir.Context.current) + except Exception: + _cur_ctx = _birth + if _birth != _cur_ctx: + raise DSLRuntimeError( + "this value was produced by a previous @jit " + "compilation and cannot be reused here: its " + "backing IR lives in an already-finalized " + "compilation context. Pass it through the " + "kernel's arguments or recompute it in this " + "compilation." + ) + # Staged-read choke: a wrapper paired with a live place cell + # materialises the CELL's current content, never a stale raw. + refreshed = _pyir_refresh_cell_read(self) + res = self.value if refreshed is None else refreshed.value else: raise ValueError( f"cannot convert {type(self)} to {dtype}, " @@ -1370,17 +1558,40 @@ def __dsl_bool__( return self.__ne__(zero, loc=loc, ip=ip) def __bool__(self) -> bool: - if isinstance(self.value, (int, float, bool)): - return bool(self.value) + v = self.__dict__.get("value") + if isinstance(v, (int, float, bool)): + # A truth-test on the payload folds trace-time data: witness it. + _pyir_record_external_payload_consumption(self) + return bool(v) else: raise DSLUserCodeError(DiagId.PHASE_DYNAMIC_TO_STATIC_BOOL) def __index__(self) -> int: - if isinstance(self.value, (int, float, bool)): - return self.value # type: ignore[return-value] + v = self.__dict__.get("value") + if isinstance(v, (int, float, bool)): + # An index coercion of the payload is a structural bake: witness it. + _pyir_record_external_payload_consumption(self) + return v # type: ignore[return-value] else: raise DSLUserCodeError(DiagId.PHASE_DYNAMIC_INDEX) + def __complex__(self) -> complex: + # The arm is a PyIR-mode fact: without PyIR there is no __complex__, + # and CPython's constructor falls through to the __index__ slot + # (``operator.index`` applies the same non-int return check). + if not is_pyir_enabled(): + return complex(operator.index(self)) + v = self.__dict__.get("value") + if isinstance(v, (int, float, bool)): + # complex() consumes the trace-time payload: witness it. + _pyir_record_external_payload_consumption(self) + return complex(v) + raise DSLUserCodeError( + DiagId.PHASE_NUMERIC_PROTOCOL_ON_STAGED, + proto="complex()", + what=type(self).__name__, + ) + def __neg__( self, *, @@ -1388,7 +1599,9 @@ def __neg__( ip: Optional[ir.InsertionPoint] = None, ) -> "Numeric": if isinstance(self.value, (bool, int, float)): - return type(self)(-self.value) + res = type(self)(-self.value) + _pyir_mark_fold_srcs(res, (self,)) + return res else: return type(self)(-self.value, loc=loc, ip=ip) # type: ignore[operator] @@ -1399,7 +1612,9 @@ def __abs__( ip: Optional[ir.InsertionPoint] = None, ) -> "Numeric": if isinstance(self.value, (bool, int, float)): - return type(self)(abs(self.value)) + res = type(self)(abs(self.value)) + _pyir_mark_fold_srcs(res, (self,)) + return res else: return type(self)(abs(self.value), loc=loc, ip=ip) # type: ignore[arg-type] @@ -1421,8 +1636,15 @@ def _from_python_value( elif isinstance(value, ArithValue): res_type = Numeric.from_mlir_type(value.type) else: + if not is_pyir_enabled(): + raise ValueError( + f"unable to convert {value} in type {type(value)} to Numeric" + ) + # Pure error path under PyIR: rendering the VALUE could emit IR (a + # traced wrapper's __str__ builds ops), so the message names the + # type only. raise ValueError( - f"unable to convert {value} in type {type(value)} to Numeric" + f"unable to convert value of type {type(value)} to Numeric" ) return res_type(value) @@ -1622,11 +1844,63 @@ def __ge__( def __pow__( self, other: Union[int, float, bool, "Numeric"], + mod: Union[int, float, bool, "Numeric", None] = None, *, loc: Optional[ir.Location] = None, ip: Optional[ir.InsertionPoint] = None, ) -> "Numeric": - return _binary_op(operator.pow)(self, other, loc=loc, ip=ip) + if mod is None: + return _binary_op(operator.pow)(self, other, loc=loc, ip=ip) + # Ternary pow (LangRef 3.3.8): payload-only fold; no staged + # modular-exponentiation form. The mod arm is a PyIR-mode fact: + # without PyIR the arm declines and Python raises its native + # TypeError, as with the two-argument signature it replaced. + if not is_pyir_enabled(): + return NotImplemented + _pyir_numeric_protocol_gate("pow(a, b, mod)", self, other, mod) + res_val = pow( + _pyir_numeric_payload(self), + _pyir_numeric_payload(other), + _pyir_numeric_payload(mod), + ) + res = type(self)(res_val) + _pyir_mark_fold_srcs(res, (self, other, mod)) + return res + + @dsl_user_op + def __rpow__( + self, + other: Union[int, float, bool, "Numeric"], + *, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, + ) -> "Numeric": + # The reflected arm is a PyIR-mode fact: without PyIR there is no + # reflected producer, so Python raises its native TypeError. + if not is_pyir_enabled(): + return NotImplemented + return _binary_op(operator.pow, flip=True)(self, other, loc=loc, ip=ip) + + def __divmod__(self, other: Any) -> Any: + # LangRef 3.3.8: divmod(a, b) is definitionally the (a // b, a % b) + # pair; each half folds or emits through its existing arm, so the + # pair inherits exactly those arms' semantics and refusals. The + # producer is a PyIR-mode fact: without PyIR the arm declines and + # Python raises its native TypeError, exactly as before. + if not is_pyir_enabled(): + return NotImplemented + if not isinstance(other, (Numeric, ArithValue, bool, int, float)): + return NotImplemented + return (self.__floordiv__(other), self.__mod__(other)) + + def __rdivmod__(self, other: Any) -> Any: + # Reflected pair (other // self, other % self) through the existing + # reflected arms; a PyIR-mode fact like the direct arm. + if not is_pyir_enabled(): + return NotImplemented + if not isinstance(other, (ArithValue, bool, int, float)): + return NotImplemented + return (self.__rfloordiv__(other), self.__rmod__(other)) def __c_pointers__(self) -> list[ctypes.c_void_p]: raise ValueError( @@ -1677,6 +1951,82 @@ def from_mlir_type(mlir_type: ir.Type) -> Type["Numeric"]: return type_map[mlir_type] +# The wrapper write funnel is INSTALLED, not merely gated. A ``property`` and a +# Python ``__setattr__`` bypass CPython's specialized attribute opcodes, so +# binding them on ``Numeric`` up front costs every construction and every +# payload read even when PyIR never runs -- and marshalling one scalar kernel +# argument is exactly that work. They are bound on first PyIR enablement +# instead. Until then ``value`` is a plain instance attribute; the payload +# lives under the same key either way, so a wrapper minted before installation +# keeps reading back correctly afterwards. + + +def _numeric_payload_get(self: "Numeric") -> Any: + """The wrapped payload; a plain-payload read by non-DSL code is a + recorded structural consumption (the coercion-witness recorder).""" + try: + v = self.__dict__["value"] + except KeyError: + raise AttributeError("value") from None + if type(v) in (bool, int, float): + _pyir_record_external_payload_consumption(self) + return v + + +def _numeric_payload_set(self: "Numeric", v: Any) -> None: + self.__dict__["value"] = v + _pyir_note_holder_write(self) + + +# The payload lives in the instance dict under its original key so every +# __dict__-driven walk sees the same storage shape as a plain attribute. +_numeric_payload_property = property(_numeric_payload_get, _numeric_payload_set) + + +# Every language-level attribute write/delete lands in plain storage (verbatim +# delegate) and stamps the write clock. +def _numeric_funnel_setattr(self: "Numeric", name: str, value: Any) -> None: + object.__setattr__(self, name, value) + _pyir_note_holder_write(self) + + +def _numeric_funnel_delattr(self: "Numeric", name: str) -> None: + object.__delattr__(self, name) + _pyir_note_holder_write(self) + + +def _pyir_install_numeric_write_funnel() -> None: + """Bind the wrapper write funnel onto ``Numeric`` and anchor the declared + fact on its now-verbatim members. + + Idempotent. ``_pyir_register_mode_fact`` calls this the first time any DSL + registers with PyIR enabled, which happens when that DSL's environment + manager is constructed -- before any tracing. + """ + if Numeric.__dict__.get("value") is _numeric_payload_property: + return + Numeric.value = _numeric_payload_property # type: ignore[assignment] + Numeric.__setattr__ = _numeric_funnel_setattr # type: ignore[assignment] + Numeric.__delattr__ = _numeric_funnel_delattr # type: ignore[assignment] + _pyir_declare_write_funnel(Numeric) + + +def _numeric_construct_read(x: "Numeric") -> Any: + """Read *x*'s current backing for wrapper-from-wrapper construction. + + Constructing a Numeric FROM another Numeric is a READ of the source at the + construction point: the same staged-read choke ``Numeric.to(ir.Value)`` applies. + + A literal-backed wrapper read inside staged CF is the place's first + in-region observation: the choke promotes it so the construction consumes + the region's carry, not a re-baked loop-invariant constant. + """ + refreshed = _pyir_refresh_cell_read(x) + if refreshed is not None: + return refreshed.value + return x.value + + def as_numeric(obj: Union[bool, int, float, ir.Value, Numeric]) -> Numeric: """Convert a Python primitive value to a Numeric type. @@ -1702,6 +2052,41 @@ def as_numeric(obj: Union[bool, int, float, ir.Value, Numeric]) -> Numeric: return Numeric._from_python_value(obj) +# ``np.array(x)`` holds a Python int in int64 up to 2**63-1 and in uint64 up to +# 2**64-1; anything wider becomes an object array whose cast to an integer dtype +# raises. Frozen here so the numpy-free cast rejects the same values. +_C_INTEGER_MIN = -(1 << 63) +_C_INTEGER_MAX = (1 << 64) - 1 + + +def _wrap_to_exact_range(value: int, exact_range: tuple[int, int]) -> int: + """Wrap ``value`` into ``exact_range`` the way a C integer cast would. + + Keeps the low ``width`` bits and reinterprets them with the range's + signedness, e.g. ``Int32(1 << 34) -> 0``. ``(0, 1)`` is the boolean range, + where every nonzero value folds to 1 instead of wrapping. + """ + if exact_range == (0, 1): + return int(bool(value)) + lo, hi = exact_range + return (value - lo) % (hi - lo + 1) + lo + + +# Bound by the DSL package initializer to the ``NumericCast`` class of the +# native extension. Read at call time, never captured at import time: the +# binding happens after this module has been imported. +_native_numeric_cast: Optional[Any] = None + + +def _float_to_int( + x: float, cast_key: Optional[tuple[int, bool]], exact_range: tuple[int, int] +) -> int: + """Narrow ``x`` to a C integer the way NumPy's scalar cast does.""" + if _native_numeric_cast is not None and cast_key is not None: + return _native_numeric_cast.float_to_int(x, *cast_key) + return _wrap_to_exact_range(int(x), exact_range) + + class Integer(Numeric, metaclass=IntegerMeta, mlir_type=T.i32, is_abstract=True): """A class representing integer values with specific width and signedness. @@ -1715,17 +2100,26 @@ class Integer(Numeric, metaclass=IntegerMeta, mlir_type=T.i32, is_abstract=True) :return: A new Integer instance with the converted value :rtype: Integer - :raises AssertionError: If the type's numpy_dtype is None + :raises AssertionError: If the type has no exactly representable range :raises NotImplementedError: If converting between different Integer types :raises ValueError: If the input type is not supported for conversion - :raises OverflowError: If converting float infinity to integer + :raises OverflowError: If converting float infinity to integer, or an + integer too wide for any C integer type Type conversion behavior: * Python scalars (bool, int, float): - * Converted through numpy dtype casting + * Converted through the target dtype's C cast * NaN and infinity values are rejected - * Example: Int8(256) -> -256 (overflow behavior) + * A value whose magnitude exceeds the target width is narrowed by + the dtype's C cast and emits ``TYPE_INT_LITERAL_OUT_OF_RANGE`` for + an integer literal (``Int8(256) -> 0``) or + ``TYPE_FLOAT_TO_INT_OUT_OF_RANGE`` for a float one + (``Int32(1e40)``). An integer wraps, dropping the high bits; a + float follows the host CPU's own conversion and is therefore + architecture-defined, matching NumPy on the same machine. To + materialize a specific bit pattern intentionally, mask first + (``Int8(256 & 0xFF)``). * MLIR Value with IntegerType: * Width differences handled by signless to signed/unsigned conversion @@ -1737,7 +2131,7 @@ class Integer(Numeric, metaclass=IntegerMeta, mlir_type=T.i32, is_abstract=True) * Example: f32 -> i32/ui32 depending on target type * Integer: - * Uses MLIR float-to-int conversion or numpy dtype casting + * Uses MLIR int-to-int conversion or the target dtype's C cast * Example: Int32(Int32(5)) => 5 * Float: @@ -1759,6 +2153,7 @@ class Integer(Numeric, metaclass=IntegerMeta, mlir_type=T.i32, is_abstract=True) # Injected by IntegerMeta.__new__ on every concrete subclass. signed: ClassVar[bool] _exact_range: ClassVar[Optional[tuple[int, int]]] + _cast_key: ClassVar[Optional[tuple[int, bool]]] def __init__( self, @@ -1768,14 +2163,14 @@ def __init__( ip: Optional[ir.InsertionPoint] = None, ) -> None: ty = type(self) - # D1: if x is a _WatchedM, ir_value() emits arith.constant AND records - # the leaf under its slot so a later mutation can rewrite the use via - # replaceAllUsesWith. + # D1: a watched meta unwraps through the declared promotion (its + # numeric class survives; ir_value() inside records the leaf for the + # retroactive rewrite) -- never to a raw signless value. if isinstance(x, _WatchedM): - x = x.ir_value() + x = as_numeric(x) if isinstance(x, (bool, int, float)): - # Add check for NaN before numpy conversion + # Add check for NaN before the cast if isinstance(x, float): if math.isnan(x): raise ValueError("Cannot convert float NaN to integer") @@ -1786,18 +2181,87 @@ def __init__( if exact_range is not None and exact_range[0] <= x <= exact_range[1]: # Already representable, so the cast below would round-trip the # value unchanged. int() truncates floats toward zero and folds - # bools to 0/1, which is what astype does for this range too. + # bools to 0/1, which is what the cast does for this range too. x_val = int(x) else: - # Out of range: defer to numpy for the wrap-around (or the - # overflow it raises on values wider than the dtype). _exact_range - # is None exactly when there is no dtype to cast to, so this is - # still where an unsupported width is rejected. - np_dtype = ty.numpy_dtype - assert np_dtype is not None, f"expects numpy.dtype, but got {np_dtype}" - x_val = int(np.array(x).astype(np_dtype)) + # Out of range: narrow the way the dtype's C cast would, and + # reject the ints too wide for any C integer. _exact_range is + # None exactly when there is no dtype to cast to, so this is + # still where an unsupported width is rejected. A float cast is + # architecture-defined and an integer cast is not, so they take + # different paths. + assert exact_range is not None, ( + f"expects an exactly representable range, but got {exact_range}" + ) + if isinstance(x, float): + x_val = _float_to_int(x, ty._cast_key, exact_range) + else: + if not (_C_INTEGER_MIN <= x <= _C_INTEGER_MAX): + raise OverflowError(f"{int(x)} is too large for a C integer") + x_val = _wrap_to_exact_range(int(x), exact_range) + # A value whose truncation lands outside the target type's width is + # silently narrowed by the cast above, losing magnitude (e.g. + # ``Int32(1 << 34) -> 0``). Surface that loss as a warning, under a + # dedicated code per literal kind so the wording and the suggested + # fix match what was written. MLIR-value / same-type paths never + # reach here, so intentional bit-pattern materialization via + # ``arith_helper.const`` is unaffected. ``Boolean`` (width 1) is + # excluded: it has no magnitude range (any nonzero is True), so the + # ``[min, max]`` formula does not apply. + # + # The two literal kinds get different bounds, because only one of + # them has a bit-pattern idiom. + # + # Integer literals are tested against the union of the signed and + # unsigned ranges of that width, and signedness is deliberately not + # part of the test. A mask, flag word or other bit pattern is + # naturally unsigned, so a full 32-bit mask reaches a signed type as + # ``Int32(0xFFFFFFFF)``; the mirror idiom ``Uint32(-1)`` spells the + # same bits the other way. Both keep all ``width`` bits -- only the + # interpretation of the sign bit changes -- and the diagnostic's own + # remedy (mask to the type width) cannot quiet them. + # + # Float literals are tested against the type's own ``[min, max]``. + # Nobody spells a bit pattern as a float, so the wider band would + # only swallow real mistakes: ``Int32(3e9)`` is out of range for the + # type and is reported, matching the diagnostic NumPy used to raise + # for that cast. Testing the truncated magnitude keeps ordinary + # fractional truncation (``Int32(3.7)``) silent -- it loses no bits. + if ty.width > 1: + int_val = int(x) + if isinstance(x, float): + lossless_lo, lossless_hi = ty.min, ty.max + else: + lossless_lo = -(1 << (ty.width - 1)) + lossless_hi = (1 << ty.width) - 1 + if int_val < lossless_lo or int_val > lossless_hi: + from .diagnostics import WarnId, report_warning + + if isinstance(x, float): + report_warning( + WarnId.TYPE_FLOAT_TO_INT_OUT_OF_RANGE, + stacklevel=3, + value=x, + type=ty.__name__, + min=ty.min, + max=ty.max, + result=x_val, + ) + else: + report_warning( + WarnId.TYPE_INT_LITERAL_OUT_OF_RANGE, + stacklevel=3, + value=int_val, + type=ty.__name__, + min=ty.min, + max=ty.max, + wrapped=x_val, + mask=(1 << ty.width) - 1, + ) elif type(x) == ty: - x_val = x.value # type: ignore[assignment] + # Same-type copy is a READ of *x* at this point: route through + # the staged-read choke, never the pinned raw. + x_val = _numeric_construct_read(x) # type: ignore[assignment] elif isinstance(x, ir.Value): x_val = x if isinstance(x.type, ir.IntegerType): @@ -1808,21 +2272,35 @@ def __init__( # float -> (u)int x_val = arith_helper.fptoi(x, ty.signed, ty.mlir_type, loc=loc, ip=ip) elif isinstance(x, Integer): + # Consult the staged-read choke BEFORE dispatching: a loop-carried + # wrapper still caches its pre-loop literal, so a ``x.value`` type + # test would bake the stale constant instead of loading the cell. + refreshed = _pyir_refresh_cell_read(x) + if refreshed is not None: + x = refreshed if isinstance(x.value, ir.Value): x_val = arith_helper.int_to_int(x.ir_value(), ty) else: - # For non-MLIR values, use numpy casting - src_val = np.array(x.value, dtype=type(x).numpy_dtype) - x_val = int(src_val.astype(ty.numpy_dtype)) + # For non-MLIR values, wrap the way the target's C cast would. + # A target with no exactly representable range (Int4, Int128) + # has no cast to perform, so the value carries across as is. + target_range = ty._exact_range + src_val = tcast(Union[bool, int, float], x.value) + x_val = ( + int(src_val) + if target_range is None + else _wrap_to_exact_range(int(src_val), target_range) + ) elif isinstance(x, Float): # float -> int is handled by Integer.__init__ recursively - Integer.__init__(self, x.value) + Integer.__init__(self, _numeric_construct_read(x)) return else: raise DSLRuntimeError(f"{x} to integer conversion is not supported") super().__init__(x_val) + @dsl_user_op def __invert__( self, *, @@ -1832,6 +2310,7 @@ def __invert__( res_type = type(self) return res_type(self.ir_value(loc=loc, ip=ip).__invert__(loc=loc, ip=ip)) + @dsl_user_op def __lshift__( self, other: Union[int, float, bool, "Numeric"], @@ -1853,6 +2332,7 @@ def __rlshift__( raise ValueError(f"Cannot left shift {other_} with {self}") return other_.__lshift__(self, loc=loc, ip=ip) # type: ignore[call-arg] + @dsl_user_op def __rshift__( self, other: Union[int, float, bool, "Numeric"], @@ -1874,6 +2354,7 @@ def __rrshift__( raise ValueError(f"Cannot right shift {other_} with {self}") return other_.__rshift__(self, loc=loc, ip=ip) # type: ignore[call-arg] + @dsl_user_op def __and__( self, other: Union[int, float, bool, "Numeric"], @@ -1892,6 +2373,7 @@ def __rand__( ) -> "Numeric": return self.__and__(other, loc=loc, ip=ip) # type: ignore[call-arg] + @dsl_user_op def __or__( self, other: Union[int, float, bool, "Numeric"], @@ -1910,6 +2392,7 @@ def __ror__( ) -> "Numeric": return self.__or__(other, loc=loc, ip=ip) # type: ignore[call-arg] + @dsl_user_op def __xor__( self, other: Union[int, float, bool, "Numeric"], @@ -1932,6 +2415,19 @@ def __tvm_ffi_int__(self) -> Union[int, ir.Value]: return self.value +# Standard IEEE-754 binary formats keyed by DSL type name, mapping to +# (``struct`` format code, max finite magnitude, smallest positive subnormal). +# Used by ``Float.__init__`` to detect a Python-float literal that overflows to +# +/-inf or underflows to 0 in the target type -- numpy-free, via exact +# ``struct`` round-trips. Only these standard formats are probed; non-IEEE / +# narrow types (bf16, tf32, fp8, fp6, fp4) are absent and narrow later in IR. +_IEEE_FLOAT_PROBE: dict = { + "Float16": ("e", 65504.0, 2.0**-24), + "Float32": ("f", 3.4028234663852886e38, 2.0**-149), + "Float64": ("d", 1.7976931348623157e308, 2.0**-1074), +} + + class Float(Numeric, metaclass=FloatMeta, mlir_type=T.f32, is_abstract=True): """A class representing floating-point values. @@ -1941,8 +2437,14 @@ class Float(Numeric, metaclass=FloatMeta, mlir_type=T.f32, is_abstract=True): Type conversion behavior: 1. Python scalars (bool, int, float): - - Converted through numpy dtype casting + + - Kept at full Python-float precision and narrowed during IR emission - Example: Float32(1.7) -> 1.7 + - A value whose magnitude cannot be represented in the target type + collapses to +/-inf (overflow) or 0 (underflow) and emits a + ``TYPE_FLOAT_LITERAL_OVERFLOW`` / ``TYPE_FLOAT_LITERAL_UNDERFLOW`` + warning, e.g. ``Float32(1e40) -> inf``. Ordinary precision/rounding + loss (``Float32(0.1)``) is not flagged. 2. MLIR Value with FloatType: - If width differs: converts between float types @@ -1978,7 +2480,6 @@ class Float(Numeric, metaclass=FloatMeta, mlir_type=T.f32, is_abstract=True): Narrow precision types and special floating-point formats support matrix on device: - :raises AssertionError: If the type's numpy_dtype is None :raises ValueError: If conversion from the input type is not supported """ @@ -1990,24 +2491,60 @@ def __init__( ip: Optional[ir.InsertionPoint] = None, ) -> None: ty = type(self) - # D1: if x is a _WatchedM, ir_value() emits arith.constant AND records - # the leaf under its slot so a later mutation can rewrite the use via - # replaceAllUsesWith. Watched Python integers use their Python signedness - # to select the integer-to-float conversion before the wrapper is erased. + # D1: a watched meta unwraps through the declared promotion (its + # numeric class survives; ir_value() inside records the leaf for the + # retroactive rewrite) -- never to a raw signless value. if isinstance(x, _WatchedM): - watched = x - x = watched.ir_value() - if isinstance(watched.python_value, bool): - x = Boolean(x) - elif isinstance(watched.python_value, int): - x = Int32(x) + x = as_numeric(x) if isinstance(x, (bool, int, float)): - # Why we need to convert x to with numpy? - # np_dtype = ty.numpy_dtype - # assert np_dtype is not None, f"expects numpy.dtype, but got {np_dtype}" - # x = float(np.array(x).astype(np_dtype)) - super().__init__(float(x)) + fx = float(x) + # A finite, nonzero Python float whose magnitude cannot be + # represented in the target type collapses to +/-inf (overflow) or + # 0 (underflow) once narrowed, losing the value entirely. Surface + # that catastrophic loss -- but NOT ordinary precision/rounding loss + # (e.g. ``Float32(0.1)``), which is inherent to every float literal + # and would be unbearably noisy. We keep the full-precision Python + # double below; the probe here is only to detect the collapse. + # + # Detection uses stdlib ``struct`` (exact IEEE-754 round-trips) so + # this stays numpy-free. Only the standard binary16/32/64 formats + # have a ``struct`` code; non-IEEE / narrow types (bf16, tf32, fp8, + # fp6, fp4) are absent from the map and narrow later during IR + # emission, so they are skipped. + struct_code, max_finite, min_subnormal = _IEEE_FLOAT_PROBE.get( + ty.__name__, (None, None, None) + ) + if struct_code is not None and fx != 0.0 and math.isfinite(fx): + try: + narrowed = struct.unpack(struct_code, struct.pack(struct_code, fx))[ + 0 + ] + except OverflowError: + # binary16 pack raises rather than saturating to inf. + narrowed = math.copysign(math.inf, fx) + if math.isinf(narrowed) or narrowed == 0.0: + from .diagnostics import WarnId, report_warning + + if math.isinf(narrowed): + report_warning( + WarnId.TYPE_FLOAT_LITERAL_OVERFLOW, + stacklevel=3, + value=fx, + type=ty.__name__, + max=max_finite, + wrapped=narrowed, + ) + else: + report_warning( + WarnId.TYPE_FLOAT_LITERAL_UNDERFLOW, + stacklevel=3, + value=fx, + type=ty.__name__, + tiny=min_subnormal, + wrapped=narrowed, + ) + super().__init__(fx) elif isinstance(x, ir.Value): if isinstance(x.type, ir.IntegerType): raise DSLRuntimeError("signless to float conversion is not implemented") @@ -2016,15 +2553,25 @@ def __init__( x = arith_helper.cvtf(x, ty.mlir_type, loc=loc, ip=ip) super().__init__(x) elif isinstance(x, Integer): + # Consult the staged-read choke BEFORE dispatching: a loop-carried + # wrapper still caches its pre-loop literal, so a ``x.value`` type + # test would bake the stale constant instead of loading the cell. + refreshed = _pyir_refresh_cell_read(x) + if refreshed is not None: + x = refreshed if isinstance(x.value, ir.Value): x = arith_helper.itofp( - x.value, type(x).signed, ty.mlir_type, loc=loc, ip=ip + x.value, + type(x).signed, + ty.mlir_type, + loc=loc, + ip=ip, ) else: x = float(x.value) # type: ignore[arg-type] super().__init__(x) elif isinstance(x, Float): - Float.__init__(self, x.value) + Float.__init__(self, _numeric_construct_read(x)) else: raise DSLRuntimeError(f"{x} to Float conversion is not supported") @@ -2071,16 +2618,16 @@ def __init__( loc: Optional[ir.Location] = None, ip: Optional[ir.InsertionPoint] = None, ) -> None: - # D1: if a is a _WatchedM, ir_value() emits arith.constant AND records - # the leaf under its slot so a later mutation can rewrite the use via - # replaceAllUsesWith. + # D1: a watched meta unwraps through the declared promotion (its + # numeric class survives; ir_value() inside records the leaf for the + # retroactive rewrite) -- never to a raw signless value. if isinstance(a, _WatchedM): - a = a.ir_value() + a = as_numeric(a) value = None if isinstance(a, (bool, int, float)): value = bool(a) elif isinstance(a, Numeric): - Boolean.__init__(self, a.value, loc=loc, ip=ip) + Boolean.__init__(self, _numeric_construct_read(a), loc=loc, ip=ip) return elif isinstance(a, ArithValue): if a.type == T.bool(): @@ -2211,14 +2758,28 @@ def __c_pointers__(self) -> list[ctypes.c_void_p]: class Float16(Float, metaclass=FloatMeta, width=16, mlir_type=T.f16): @staticmethod def _get_c_pointer(value: float) -> ctypes.c_void_p: - # Convert float to float16 binary representation - # First convert to numpy float16 to handle the conversion - f16_val = np.float16(value) - # Get the raw bits as a 16-bit integer - bits: int = int(f16_val.view(np.uint16)) - # Create a short (16-bit int) with those bits - c_val = ctypes.c_short(int(bits)) - return _make_owning_c_pointer(c_val) + """Marshal ``value`` as its IEEE-754 binary16 bit pattern. + + Two cases need handling beyond a plain ``struct`` pack: + + * NaN. ``struct`` collapses every NaN to the canonical quiet pattern, + which would turn a signaling NaN into a quiet one and drop the + payload. Narrow the payload explicitly instead, keeping the high + mantissa bits and forcing a nonzero payload so the result cannot + decay into an infinity. + * Finite overflow. ``struct`` raises rather than saturating to inf. + """ + if value != value: + double_bits = struct.unpack("> 48) & 0x8000 + payload = (double_bits & ((1 << 52) - 1)) >> 42 + bits = sign | 0x7C00 | (payload or 1) + else: + try: + bits = struct.unpack(" list[ctypes.c_void_p]: if not isinstance(self.value, float): @@ -2232,13 +2793,15 @@ def __c_pointers__(self) -> list[ctypes.c_void_p]: raise ValueError("only float is supported") # Convert float32 to bfloat16 representation # First convert the value to float32 bit representation - f32_val = np.float32(self.value) - # Get the 32-bit integer representation - bits = int(f32_val.view(np.uint32)) + try: + bits = struct.unpack("> 16) + bf16_bits = bits >> 16 # Create a short (16-bit int) with those bits - c_val = ctypes.c_short(bf16_bits) # type: ignore[arg-type] + c_val = ctypes.c_short(bf16_bits) c_pointer = _make_owning_c_pointer(c_val) return [c_pointer] @@ -3382,16 +3945,17 @@ def implicitDowncastNumericType( ) -> Union[bool, int, float, ir.Value]: if isinstance(value, Numeric): return value.ir_value() - # D1: ``_WatchedM`` is a meta-promotion wrapper used by PyIR to track - # primitive reads inside staged CF. When passed to a constant-emitting - # callsite, bake its IR (which records the leaf under the slot so a - # later mutation can rewrite it via pyir.load %ref). + # A direct ``_mlir.dialects`` builder call receives the plain payload of a + # watched-META wrapper (OFF-parity: the un-watched value is the payload); + # the bake is a recorded non-retargetable consumption, and a promoted + # place refuses (no single compile-time value exists to pass). try: from .pyir_runtime import _WatchedM as _PyIRWatchedM + from .pyir_runtime import _pyir_boundary_consume_meta_arg as _pyir_consume except ImportError: _PyIRWatchedM = None # type: ignore[misc,assignment] if _PyIRWatchedM is not None and isinstance(value, _PyIRWatchedM): - return value.ir_value() + return _pyir_consume(value) return value diff --git a/python/CuTeDSL/cutlass/base_dsl/utils/__init__.py b/python/CuTeDSL/cutlass/base_dsl/utils/__init__.py index b2458957d7..f4a3a99c04 100644 --- a/python/CuTeDSL/cutlass/base_dsl/utils/__init__.py +++ b/python/CuTeDSL/cutlass/base_dsl/utils/__init__.py @@ -12,6 +12,7 @@ from . import stacktrace from . import logger from . import timer + __all__ = [ "logger", "timer", diff --git a/python/CuTeDSL/cutlass/base_dsl/utils/leaf_utils.py b/python/CuTeDSL/cutlass/base_dsl/utils/leaf_utils.py index 2e1ec7253b..47ee83c3a3 100644 --- a/python/CuTeDSL/cutlass/base_dsl/utils/leaf_utils.py +++ b/python/CuTeDSL/cutlass/base_dsl/utils/leaf_utils.py @@ -80,6 +80,33 @@ def _is_assignable_leaf(obj: Any) -> bool: return False +def _has_assignable_leaves(obj: Any, visited: set[int]) -> bool: + """Read-only reachability query: does obj's subtree hold any assignable leaf?""" + if obj is None: + return False + if _is_assignable_leaf(obj): + return True + if id(obj) in visited: + return False + visited.add(id(obj)) + if isinstance(obj, (list, tuple)): + return any(_has_assignable_leaves(v, visited) for v in obj) + if isinstance(obj, dict): + return any(_has_assignable_leaves(v, visited) for v in obj.values()) + if is_frozen_dataclass(obj): + return any( + _has_assignable_leaves(getattr(obj, f.name), visited) + for f in dataclasses.fields(obj) + ) + if hasattr(obj, "__dict__") or hasattr(type(obj), "__slots__"): + return any( + _has_assignable_leaves(v, visited) + for k, v in _get_all_attrs(obj).items() + if not k.startswith("__") and not callable(v) + ) + return False + + def _flatten_to_ir_values(values_dict: Any) -> list[ir.Value]: """Flatten a values_dict from __extract_mlir_values__ to list of ir.Values.""" result = [] @@ -306,6 +333,13 @@ def gather_leaves( immutable_proxies: list[tuple[Any, ...]] = [] visited = set() + # A proxy may only become VISIBLE inside framework-owned containers (the + # capture list itself or another proxy of this walk): a real user + # container keeps its original object between gather and inject, so a + # caller that never injects -- or user code running inside the window -- + # never observes a half-disassembled object. inject_leaves' Phase-2 + # reconstruction writes the rebuilt immutable into the real parent. + framework_parents = {id(objects)} def _gather_recursive( obj: Any, parent: Any, key: Any, key_type: str, path: str @@ -348,7 +382,9 @@ def _gather_recursive( if isinstance(obj, tuple): proxy_list = list(obj) immutable_proxies.append((obj, proxy_list, parent, key, key_type)) - _write_into_parent(parent, key, key_type, proxy_list) + if id(parent) in framework_parents: + _write_into_parent(parent, key, key_type, proxy_list) + framework_parents.add(id(proxy_list)) for i, item in enumerate(proxy_list): item_path = f"{path}({i})" if path else f"({i})" _gather_recursive(item, proxy_list, i, "list", item_path) @@ -356,10 +392,17 @@ def _gather_recursive( # Frozen dataclass (IMMUTABLE container -- create mutable proxy) if is_frozen_dataclass(obj): + # Stateless frozen dataclasses (no ir.Value-carrying leaves anywhere + # below) have nothing to gather or inject; installing a proxy would + # only leak a SimpleNamespace into the caller's object. + if not _has_assignable_leaves(obj, set()): + return fields = dataclasses.fields(obj) proxy = SimpleNamespace(**{f.name: getattr(obj, f.name) for f in fields}) immutable_proxies.append((obj, proxy, parent, key, key_type)) - _write_into_parent(parent, key, key_type, proxy) + if id(parent) in framework_parents: + _write_into_parent(parent, key, key_type, proxy) + framework_parents.add(id(proxy)) for f in fields: attr_val = getattr(proxy, f.name) if f.name.startswith("__") or callable(attr_val): diff --git a/python/CuTeDSL/cutlass/base_dsl/utils/mlir_value_tree.py b/python/CuTeDSL/cutlass/base_dsl/utils/mlir_value_tree.py index d02ddbfe83..840c5b5889 100644 --- a/python/CuTeDSL/cutlass/base_dsl/utils/mlir_value_tree.py +++ b/python/CuTeDSL/cutlass/base_dsl/utils/mlir_value_tree.py @@ -21,9 +21,10 @@ orchestration. """ +import contextlib import dataclasses from types import SimpleNamespace -from typing import Any, get_origin +from typing import Any, Iterator, get_origin from ..common import DSLUserCodeError from ..diagnostics import DiagId @@ -34,6 +35,19 @@ _unflatten_mlir_values, ) from ..._mlir import ir +from ..pyir_state import _EXTRACTION_WALK_OWNERS + + +@contextlib.contextmanager +def _extraction_walk_owner(obj: Any) -> Iterator[None]: + """Declare *obj* as the extraction walk's current owner context: a leaf + materialization under the walk resolves its place row against the declared + holders, never a first-sighted registry candidate.""" + _EXTRACTION_WALK_OWNERS.append(obj) + try: + yield + finally: + _EXTRACTION_WALK_OWNERS.pop() def is_dynamic_expression(value: object) -> bool: @@ -68,21 +82,24 @@ def extract_mlir_values(obj: object, *, structured: bool = False) -> Any: if structured: # Tree-structured mode: return __extract_mlir_values__ result directly if hasattr(obj, "__extract_mlir_values__"): - return obj.__extract_mlir_values__() + with _extraction_walk_owner(obj): + return obj.__extract_mlir_values__() elif dataclasses.is_dataclass(obj) and not isinstance(obj, type): - return { - field.name: extract_mlir_values( - getattr(obj, field.name), structured=True - ) - for field in dataclasses.fields(obj) - } + with _extraction_walk_owner(obj): + return { + field.name: extract_mlir_values( + getattr(obj, field.name), structured=True + ) + for field in dataclasses.fields(obj) + } elif isinstance(obj, (tuple, list)): return [extract_mlir_values(x, structured=True) for x in obj] elif isinstance(obj, SimpleNamespace): - return { - k: extract_mlir_values(v, structured=True) - for k, v in obj.__dict__.items() - } + with _extraction_walk_owner(obj): + return { + k: extract_mlir_values(v, structured=True) + for k, v in obj.__dict__.items() + } elif isinstance(obj, ir.Value): return obj elif isinstance(obj, ir.BlockArgumentList): @@ -94,13 +111,15 @@ def extract_mlir_values(obj: object, *, structured: bool = False) -> Any: res = [] if hasattr(obj, "__extract_mlir_values__"): # Flatten whatever __extract_mlir_values__ returns to ensure we always get a flat list - res = flatten_mlir_values(obj.__extract_mlir_values__()) + with _extraction_walk_owner(obj): + res = flatten_mlir_values(obj.__extract_mlir_values__()) elif isinstance(obj, (tuple, list)): res = sum((extract_mlir_values(x) for x in obj), []) elif isinstance(obj, SimpleNamespace): res = [] - for k, v in obj.__dict__.items(): - res.extend(extract_mlir_values(v)) + with _extraction_walk_owner(obj): + for k, v in obj.__dict__.items(): + res.extend(extract_mlir_values(v)) elif isinstance(obj, set): raise DSLUserCodeError( DiagId.ARG_UNORDERED_CONTAINER, @@ -193,7 +212,16 @@ def new_from_mlir_values(obj: Any, values: Any, *, structured: bool = False) -> """ # Objects with __new_from_mlir_values__ always receive values directly if hasattr(obj, "__new_from_mlir_values__"): - return obj.__new_from_mlir_values__(values) + rebuilt = obj.__new_from_mlir_values__(values) + # Re-root the rebuilt object's place at the source's owner token so + # place-keyed lookups re-resolve; guarded, token tables only. + try: + from ..pyir_runtime import _pyir_adopt_rebuilt_owner_token + + _pyir_adopt_rebuilt_owner_token(obj, rebuilt) + except Exception: + pass + return rebuilt if structured: # Tree-structured mode diff --git a/python/CuTeDSL/cutlass/base_dsl/utils/tree_utils.py b/python/CuTeDSL/cutlass/base_dsl/utils/tree_utils.py index 598c564327..2aa17dfb73 100644 --- a/python/CuTeDSL/cutlass/base_dsl/utils/tree_utils.py +++ b/python/CuTeDSL/cutlass/base_dsl/utils/tree_utils.py @@ -475,7 +475,15 @@ def dynamic_expression_from_iterable( values = _unflatten_mlir_values(children_list, metadata.template) else: values = children_list - return metadata.original_obj.__new_from_mlir_values__(values) + rebuilt = metadata.original_obj.__new_from_mlir_values__(values) + # The pytree unflatten bypasses ``new_from_mlir_values``: re-root the rebuilt + # node's place at the template's owner token here too (tables only; the + # token table pins non-weakrefable keys itself, so adoption cannot raise + # for any admitted object). + from ..pyir_runtime import _pyir_adopt_rebuilt_owner_token + + _pyir_adopt_rebuilt_owner_token(metadata.original_obj, rebuilt) + return rebuilt def default_dict_to_iterable(x: Any) -> tuple[SimpleNamespace, list[Any]]: @@ -861,6 +869,12 @@ def _tree_flatten( ) else: + # Transparency protocol (F-TRANSPARENT): an observation wrapper that + # declares its plain counterpart flattens AS its plain value -- the + # wrapper type is a host-trace artifact and never crosses a boundary. + # Protocol dispatch by getattr: no wrapper-class imports here. + if getattr(type(x), "__pyir_plain_type__", None) is not None: + return _tree_flatten(x.__pyir_plain_view__(), return_ir_values) node_type = get_registered_node_types_or_insert(x) if node_type: node_metadata, children = node_type.to_iterable(x) @@ -942,7 +956,16 @@ def _tree_unflatten(treedef: PyTreeDef | Leaf, xs: Iterator[Any]) -> Any: return None metadata = getattr(treedef, "node_metadata", None) if metadata and getattr(metadata, "is_dynamic_expression", False): - return metadata.original_obj.__new_from_mlir_values__([next(xs)]) + rebuilt = metadata.original_obj.__new_from_mlir_values__([next(xs)]) + # Pytree-leaf reconstruction of a dynamic-expression node: re-root + # its place at the template token (guarded, token tables only). + try: + from ..pyir_runtime import _pyir_adopt_rebuilt_owner_token + + _pyir_adopt_rebuilt_owner_token(metadata.original_obj, rebuilt) + except Exception: + pass + return rebuilt if getattr(treedef, "is_numeric", False): return as_numeric(next(xs)) return next(xs) diff --git a/python/CuTeDSL/cutlass/cute/__init__.py b/python/CuTeDSL/cutlass/cute/__init__.py index 064bc5bf43..e4861af701 100644 --- a/python/CuTeDSL/cutlass/cute/__init__.py +++ b/python/CuTeDSL/cutlass/cute/__init__.py @@ -227,6 +227,8 @@ from .ffi import ffi, extern, BitCode, ConstValue, mangle +from .launch_facts import _get_launch_facts + # Aliases jit: Callable[..., Any] = _dsl.CuTeDSL.jit kernel: Callable[..., Any] = _dsl.CuTeDSL.kernel diff --git a/python/CuTeDSL/cutlass/cute/algorithm.py b/python/CuTeDSL/cutlass/cute/algorithm.py index 3769b94b5d..853b6d7ba9 100644 --- a/python/CuTeDSL/cutlass/cute/algorithm.py +++ b/python/CuTeDSL/cutlass/cute/algorithm.py @@ -550,8 +550,10 @@ def copy( logical profile ``((ATOM_V,ATOM_REST),REST,...)``, the predication tensor must maintain profile compatibility with ``(ATOM_REST,REST,...)``. - For Copy Atoms requiring single-threaded execution, thread election is managed automatically by the - copy operation. External thread selection mechanisms are not necessary. + In general, Copy Atoms requiring single-threaded execution handle thread + election internally. The non-tensor bulk copy operations (``CopyBulkG2SOp``, + ``CopyBulkG2SMulticastOp``, ``CopyBulkS2GOp``, ``CopyBulkS2GByteMaskOp``, and + ``CopyBulkS2SOp``) are exceptions: they do not perform thread election. .. note:: @@ -686,8 +688,7 @@ def prefetch( rest-dimension entry, so a single ``cute.prefetch`` call covers every stage and lives outside any per-stage loop. - For Copy Atoms that require single-threaded execution, the copy op automatically handles thread - election internally. Manual thread selection is not required in such cases. + Non-tensor bulk copy prefetch requires explicit thread election. """ src_list = _normalize_variadic_tensor_operand(src, "src") diff --git a/python/CuTeDSL/cutlass/cute/arch/__init__.py b/python/CuTeDSL/cutlass/cute/arch/__init__.py index bfedb2aaf2..b0e5488782 100644 --- a/python/CuTeDSL/cutlass/cute/arch/__init__.py +++ b/python/CuTeDSL/cutlass/cute/arch/__init__.py @@ -133,6 +133,8 @@ "fmin", "rcp_approx", "exp2", + "exp2_packed_f16x2", + "exp2_packed_bf16x2", "cvt_i8x4_to_f32x4", "cvt_i8x2_to_f32x2", "cvt_i8_bf16", diff --git a/python/CuTeDSL/cutlass/cute/arch/nvvm_wrappers.py b/python/CuTeDSL/cutlass/cute/arch/nvvm_wrappers.py index e827c76ef6..59b6cfedf6 100644 --- a/python/CuTeDSL/cutlass/cute/arch/nvvm_wrappers.py +++ b/python/CuTeDSL/cutlass/cute/arch/nvvm_wrappers.py @@ -1637,6 +1637,30 @@ def _add_packed_half2_ptx( ) +def _exp2_packed_half2_ptx( + a: Int32, + *, + dtype: Type[Numeric], + predicate: Optional[Boolean], + loc: Optional[ir.Location], + ip: Optional[ir.InsertionPoint], +) -> Int32: + if dtype is Float16: + ptx = "ex2.approx.f16x2 {$w0}, {$r0};" + elif dtype is BFloat16: + ptx = "ex2.approx.ftz.bf16x2 {$w0}, {$r0};" + else: + raise TypeError("dtype must be cutlass.Float16 or cutlass.BFloat16") + return inline_ptx( + ptx, + write_only_types=[Int32], + read_only_args=[a], + predicate=predicate, + loc=loc, + ip=ip, + ) + + def _ensure_tuple2(src: Any, name: str) -> Tuple[Any, Any]: if not isinstance(src, tuple) or len(src) != 2: raise TypeError(f"{name} must be a 2-tuple") @@ -1925,6 +1949,42 @@ def _add_packed_half2( ) +def _exp2_packed_half2( + a: Union[Int32, Tuple[Numeric, Numeric]], + *, + dtype: Type[_NumericT], + predicate: Optional[Boolean], + loc: Optional[ir.Location], + ip: Optional[ir.InsertionPoint], +) -> Union[Int32, Tuple[_NumericT, _NumericT]]: + if dtype not in (Float16, BFloat16): + raise TypeError("dtype must be cutlass.Float16 or cutlass.BFloat16") + if isinstance(a, tuple): + src_a = _ensure_tuple2(a, "a") + vec_a = _tuple_to_vec(src_a, dtype, loc=loc, ip=ip) + packed_a = _vector_to_packed_half2(vec_a, loc=loc, ip=ip) + packed_res = _exp2_packed_half2_ptx( + packed_a, + dtype=dtype, + predicate=predicate, + loc=loc, + ip=ip, + ) + return _unpack_vec2( + _packed_half2_to_vector(packed_res, dtype, loc=loc, ip=ip), + dtype, + loc=loc, + ip=ip, + ) + return _exp2_packed_half2_ptx( + a, + dtype=dtype, + predicate=predicate, + loc=loc, + ip=ip, + ) + + def _sub_packed_half2( a: Union[Int32, Tuple[Numeric, Numeric]], b: Union[Int32, Tuple[Numeric, Numeric]], @@ -2189,6 +2249,48 @@ def add_packed_bf16x2( ) +@dsl_user_op +def exp2_packed_f16x2( + a: Union[Int32, Tuple[Numeric, Numeric]], + *, + predicate: Optional[Boolean] = None, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, +) -> Union[Int32, Tuple[Float16, Float16]]: + """Approximate base-2 exponential for packed f16x2 values. + + Lowers to PTX ``ex2.approx.f16x2``. + """ + return _exp2_packed_half2( + a, + dtype=Float16, + predicate=predicate, + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def exp2_packed_bf16x2( + a: Union[Int32, Tuple[Numeric, Numeric]], + *, + predicate: Optional[Boolean] = None, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, +) -> Union[Int32, Tuple[BFloat16, BFloat16]]: + """Approximate base-2 exponential for packed bf16x2 values. + + Lowers to PTX ``ex2.approx.ftz.bf16x2``. + """ + return _exp2_packed_half2( + a, + dtype=BFloat16, + predicate=predicate, + loc=loc, + ip=ip, + ) + + @dsl_user_op def fma_f16( a: Numeric, @@ -4703,7 +4805,31 @@ def convert_arg(arg: object) -> Union[ir.Value, Numeric]: return llvm_ptr return as_numeric(arg) - read_only_ir = [convert_arg(arg) for arg in read_only_args] + def convert_read_only_arg(arg: object) -> Union[ir.Value, Numeric]: + """Keep Float32 inputs in registers when lowering inline PTX. + + The 4.8 NVVM lowering can assign the integer-immediate constraint + ``n`` when a Float32 operand folds to a constant. An explicit register + move prevents that invalid constraint while preserving the operand's + value and bit pattern. + """ + converted = convert_arg(arg) + if isinstance(converted, Float32): + value = converted.ir_value(loc=loc, ip=ip) + elif isinstance(converted, ir.Value) and converted.type == Float32.mlir_type: + value = converted + else: + return converted + return llvm.inline_asm( + Float32.mlir_type, + [value], + "mov.f32 $0, $1;", + "=f,f", + loc=loc, + ip=ip, + ) + + read_only_ir = [convert_read_only_arg(arg) for arg in read_only_args] read_write_ir = [convert_arg(arg) for arg in read_write_args] # Build write_only result types diff --git a/python/CuTeDSL/cutlass/cute/arch/tmem.py b/python/CuTeDSL/cutlass/cute/arch/tmem.py index 2e091dc720..66af1b93f7 100644 --- a/python/CuTeDSL/cutlass/cute/arch/tmem.py +++ b/python/CuTeDSL/cutlass/cute/arch/tmem.py @@ -19,6 +19,7 @@ import cutlass._mlir.dialects.cute as _cute_ir import cutlass._mlir.dialects.cute_nvgpu as _cute_nvgpu_ir +from ..core import struct from ..typing import Pointer, Int, Int32, Numeric, NumericMeta SM100_TMEM_CAPACITY_COLUMNS = ( @@ -128,6 +129,9 @@ def retrieve_tmem_ptr( f"element_type must be a type of Numeric, but got {element_type}" ) + if isinstance(ptr_to_buffer_holding_addr, struct._ScalarData): + ptr_to_buffer_holding_addr = ptr_to_buffer_holding_addr.ptr + res_ty = _cute_ir.PtrType.get(element_type.mlir_type, AddressSpace.tmem, alignment) return _cute_nvgpu_ir.arch_sm100_retrieve_tmem_ptr( res_ty, @@ -178,6 +182,9 @@ def alloc_tmem( err_msg += f", but got {num_columns}." raise ValueError(err_msg) + if isinstance(smem_ptr_to_write_address, struct._ScalarData): + smem_ptr_to_write_address = smem_ptr_to_write_address.ptr + _cute_nvgpu_ir.arch_sm100_alloc_tmem( Int32(num_columns).ir_value(loc=loc, ip=ip), smem_ptr_to_write_address.value, @@ -246,6 +253,9 @@ def dealloc_tmem( err_msg += f", but got {num_columns}." raise ValueError(err_msg) + if isinstance(tmem_ptr, struct._ScalarData): + tmem_ptr = tmem_ptr.ptr + _cute_nvgpu_ir.arch_sm100_dealloc_tmem( tmem_ptr.value, Int32(num_columns).ir_value(loc=loc, ip=ip), diff --git a/python/CuTeDSL/cutlass/cute/launch_facts.py b/python/CuTeDSL/cutlass/cute/launch_facts.py new file mode 100644 index 0000000000..ca0256830a --- /dev/null +++ b/python/CuTeDSL/cutlass/cute/launch_facts.py @@ -0,0 +1,168 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: LicenseRef-NvidiaProprietary +# +# Use of this software is governed by the terms and conditions of the +# NVIDIA End User License Agreement (EULA), available at: +# https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/license.html +# +# Any use, reproduction, disclosure, or distribution of this software +# and related documentation outside the scope permitted by the EULA +# is strictly prohibited. + +"""Private provider access to verified facts about one traced launch. + +The launch backend attaches a versioned metadata dictionary to each supported +kernel operation before tracing its body. Absent optional fields mean the +backend could not prove one static value. Missing, malformed, or unsupported +metadata raises :class:`DSLRuntimeError` instead of returning partial facts. +""" + +from dataclasses import dataclass + +from cutlass.cutlass_dsl import DSLRuntimeError +from cutlass.cutlass_dsl._launch_facts_metadata import ( + CLUSTER_LAUNCH_FIELD, + COOPERATIVE_LAUNCH_FIELD, + EXACT_BLOCK_DIM_FIELD, + EXACT_CLUSTER_DIM_FIELD, + EXACT_GRID_DIM_FIELD, + LAUNCH_FACTS_ATTR, + LAUNCH_FACTS_SCHEMA_VERSION, + LAUNCH_FACTS_SCHEMA_VERSION_FIELD, +) +from cutlass._mlir import ir + + +_KERNEL_OP_NAMES = frozenset({"cuda.kernel", "lir.func"}) + + +def _current_kernel(ip: ir.InsertionPoint | None) -> ir.Operation: + """Return the enclosing kernel while a supported CuTe kernel is traced.""" + + if ip is not None: + current_ip = ip + else: + try: + current_ip = ir.InsertionPoint.current + except Exception as exc: + raise DSLRuntimeError( + "_get_launch_facts must be called while tracing a CuTe kernel" + ) from exc + + if current_ip is None or current_ip.block is None: + raise DSLRuntimeError( + "_get_launch_facts must be called while tracing a CuTe kernel" + ) + + op = current_ip.block.owner + while op is not None: + operation = getattr(op, "operation", op) + if operation.name in _KERNEL_OP_NAMES: + return operation + op = operation.parent + raise DSLRuntimeError( + "_get_launch_facts must be called while tracing a CuTe kernel" + ) + + +@dataclass(frozen=True) +class LaunchFacts: + """Statically verified topology facts for one traced kernel launch. + + Each dimension field is a positive three-element tuple or ``None`` when the + value is runtime-dynamic or not uniquely determined. In particular, + ``exact_cluster_dim`` is absent when preferred and fallback cluster shapes + differ. A launch-mode field is ``None`` when the backend cannot prove a + static Boolean value. + """ + + exact_block_dim: tuple[int, int, int] | None = None + exact_grid_dim: tuple[int, int, int] | None = None + exact_cluster_dim: tuple[int, int, int] | None = None + cooperative_launch: bool | None = None + cluster_launch: bool | None = None + + +def _required_version(attrs: ir.DictAttr) -> None: + try: + version = ir.IntegerAttr(attrs[LAUNCH_FACTS_SCHEMA_VERSION_FIELD]).value + except (KeyError, TypeError, ValueError) as exc: + raise DSLRuntimeError( + "launch facts metadata has no valid schema version" + ) from exc + if version != LAUNCH_FACTS_SCHEMA_VERSION: + raise DSLRuntimeError( + f"unsupported launch facts schema version {version}; " + f"expected {LAUNCH_FACTS_SCHEMA_VERSION}" + ) + + +def _optional_dim(attrs: ir.DictAttr, name: str) -> tuple[int, int, int] | None: + try: + raw = attrs[name] + except KeyError: + return None + try: + values = tuple(int(value) for value in ir.DenseI64ArrayAttr(raw)) + except (TypeError, ValueError) as exc: + raise DSLRuntimeError( + f"launch facts field {name!r} must be a three-dimensional integer array" + ) from exc + if len(values) != 3 or any(value <= 0 for value in values): + raise DSLRuntimeError( + f"launch facts field {name!r} must contain three positive dimensions" + ) + return values # type: ignore[return-value] + + +def _optional_bool(attrs: ir.DictAttr, name: str) -> bool | None: + try: + raw = attrs[name] + except KeyError: + return None + try: + return ir.BoolAttr(raw).value + except (TypeError, ValueError) as exc: + raise DSLRuntimeError(f"launch facts field {name!r} must be a bool") from exc + + +def _get_launch_facts( + *, + ip: ir.InsertionPoint | None = None, +) -> LaunchFacts: + """Return the static launch facts attached to the enclosing CuTe kernel. + + This private provider API is available only while tracing a ``cuda.kernel`` + or ``lir.func``. It does not expose the complete launch configuration; + fields whose values remain dynamic are returned as ``None``. + + Raises + ------ + DSLRuntimeError + If no supported kernel is being traced, launch metadata is absent, or + its dictionary, schema version, dimensions, or Boolean fields are + malformed. + """ + + kernel = _current_kernel(ip) + try: + raw_attrs = kernel.attributes[LAUNCH_FACTS_ATTR] + except KeyError as exc: + raise DSLRuntimeError( + "launch facts are unavailable on the enclosing CuTe kernel" + ) from exc + try: + attrs = ir.DictAttr(raw_attrs) + except (TypeError, ValueError) as exc: + raise DSLRuntimeError("launch facts metadata must be a dictionary") from exc + _required_version(attrs) + return LaunchFacts( + exact_block_dim=_optional_dim(attrs, EXACT_BLOCK_DIM_FIELD), + exact_grid_dim=_optional_dim(attrs, EXACT_GRID_DIM_FIELD), + exact_cluster_dim=_optional_dim(attrs, EXACT_CLUSTER_DIM_FIELD), + cooperative_launch=_optional_bool(attrs, COOPERATIVE_LAUNCH_FIELD), + cluster_launch=_optional_bool(attrs, CLUSTER_LAUNCH_FIELD), + ) + + +__all__: list[str] = [] diff --git a/python/CuTeDSL/cutlass/cute/nvgpu/cpasync/copy.py b/python/CuTeDSL/cutlass/cute/nvgpu/cpasync/copy.py index dd70c68df8..e4a5390503 100644 --- a/python/CuTeDSL/cutlass/cute/nvgpu/cpasync/copy.py +++ b/python/CuTeDSL/cutlass/cute/nvgpu/cpasync/copy.py @@ -420,7 +420,6 @@ def with_( return CopyBulkTensor2DGather4G2STrait(self.unpack(loc=loc, ip=ip, **kwargs)) - class CopyBulkTensor2DGather4G2STrait(Trait): pass @@ -585,6 +584,67 @@ class CopyBulkTensorTileG2SMulticastTrait(Trait): pass +@dataclass +class CopyBulkTensor2DGather4G2SMulticastOp(CopyG2STileBaseOp): + """ + Bulk tensor asynchronous multicast GMEM to SMEM Copy Operation using the TMA unit. + + See the `PTX documentation `__. + This Operation uses TMA in the ``.tile::gather4`` mode. + """ + + def __post_init__(self) -> None: + super().__post_init__() + # base_dsl.Arch verification + arch: base_dsl.Arch = BaseDSL._get_dsl().get_arch_enum() + if not arch >= base_dsl.Arch.sm_100: + raise DSLUserCodeError( + f"expects arch to be at least {base_dsl.Arch.sm_100.name}, but got {arch.name}", + suggestion="Ensure env CUTE_DSL_ARCH matches your GPU architecture", + ) + + def _get_description(self) -> str: + return "cp.async GMEM -> SMEM bulk tensor gather4 multicast copy Operation" + + def _make_trait( + self, + copy_internal_type: Type[Numeric], + *, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, + **kwargs: Any, + ) -> "CopyBulkTensor2DGather4G2SMulticastNonExecTrait": + raise NotImplementedError( + "Use cpasync.make_tiled_tma_atom with gmem_coord_tensor to obtain a copy Atom for TMA" + ) + + def _to_ir(self) -> _cute_nvgpu_ir.GatherScatterTmaLoadEnum: + if self.cta_group == CtaGroup.ONE: + return _cute_nvgpu_ir.GatherScatterTmaLoadEnum.sm_100_multicast + elif self.cta_group == CtaGroup.TWO: + return _cute_nvgpu_ir.GatherScatterTmaLoadEnum.sm_100_2sm_multicast + else: + assert False, "unrecognized self.cta_group" + + +class CopyBulkTensor2DGather4G2SMulticastNonExecTrait( + CopyG2STileMulticastNonExecBaseTrait +): + def with_( + self, + *, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, + **kwargs: Any, + ) -> "CopyBulkTensor2DGather4G2SMulticastTrait": + return CopyBulkTensor2DGather4G2SMulticastTrait( + self.unpack(loc=loc, ip=ip, **kwargs) + ) + + +class CopyBulkTensor2DGather4G2SMulticastTrait(Trait): + pass + @dataclass class CopyBulkTensorIm2ColG2SOp(TmaCopyOp): @@ -970,6 +1030,56 @@ class CopyBulkTensorTileS2GTrait(Trait): pass +class CopyBulkTensor2DScatter4S2GNonExecTrait(CopyBulkTensorTileS2GNonExecTrait): + def with_( + self, + *, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, + **kwargs: Any, + ) -> "CopyBulkTensor2DScatter4S2GTrait": + return CopyBulkTensor2DScatter4S2GTrait(self.unpack(loc=loc, ip=ip, **kwargs)) + + +class CopyBulkTensor2DScatter4S2GTrait(CopyBulkTensorTileS2GTrait): + pass + + +@dataclass +class CopyBulkTensor2DScatter4S2GOp(TmaCopyOp): + """ + Bulk tensor asynchronous SMEM to GMEM Copy Operation using TMA + ``.tile::scatter4`` mode. + + The destination operand to :func:`cute.copy` is zipped as + ``[dst_coord_tensor, index_tensor]``. The coord tensor supplies the common + 2D TMA coordinate and the index tensor supplies the four scatter indices. + """ + + def __post_init__(self) -> None: + arch = BaseDSL._get_dsl().get_arch_enum() + if not arch >= base_dsl.Arch.sm_100: + raise DSLUserCodeError( + f"expects arch to be at least {base_dsl.Arch.sm_100.name}, but got {arch.name}", + suggestion="Ensure env CUTE_DSL_ARCH matches your GPU architecture", + ) + + def __str__(self) -> str: + return "cp.async SMEM -> GMEM bulk tensor scatter4 copy Operation" + + def _make_trait( + self, + copy_internal_type: Type[Numeric], + *, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, + **kwargs: Any, + ) -> "CopyBulkTensor2DScatter4S2GNonExecTrait": + raise NotImplementedError( + "Use cpasync.make_tiled_tma_atom with gmem_coord_tensor to obtain a copy Atom for TMA" + ) + + @dataclass class CopyReduceBulkTensorTileS2GOp(TmaCopyOp): """ @@ -1505,12 +1615,7 @@ class CopyBulkG2SOp(CopyOp): """ Bulk copy asynchronous GMEM to SMEM Copy Operation. - Invoke :func:`cute.copy` collectively from a converged warp. The compiler - elects one issuing lane for this operation; do not wrap the call in - :func:`cute.arch.elect_one`. An outer election would leave only one lane - able to reach the compiler-generated full-warp election, creating an - invalid synchronization that can deadlock. With NVVM diagnostics enabled, - the compiler rejects this pattern. + The compiler does not perform thread election for this operation. See the `PTX documentation `__. """ @@ -1601,12 +1706,7 @@ class CopyBulkG2SMulticastOp(CopyOp): """ Bulk multicast copy asynchronous GMEM to SMEM Copy Operation. - Invoke :func:`cute.copy` collectively from a converged warp. The compiler - elects one issuing lane for this operation; do not wrap the call in - :func:`cute.arch.elect_one`. An outer election would leave only one lane - able to reach the compiler-generated full-warp election, creating an - invalid synchronization that can deadlock. With NVVM diagnostics enabled, - the compiler rejects this pattern. + The compiler does not perform thread election for this operation. See the `PTX documentation `__. """ @@ -1706,12 +1806,7 @@ class CopyBulkS2GOp(CopyOp): """ Bulk copy asynchronous SMEM to GMEM Copy Operation. - Invoke :func:`cute.copy` collectively from a converged warp. The compiler - elects one issuing lane for this operation; do not wrap the call in - :func:`cute.arch.elect_one`. An outer election would leave only one lane - able to reach the compiler-generated full-warp election, creating an - invalid synchronization that can deadlock. With NVVM diagnostics enabled, - the compiler rejects this pattern. + The compiler does not perform thread election for this operation. See the `PTX documentation `__. """ @@ -1765,12 +1860,7 @@ class CopyBulkS2GByteMaskOp(CopyOp): The i-th bit in the 16-bit wide byteMask operand specifies whether the i-th byte of each 16-byte wide chunk of source data is copied to the destination. - Invoke :func:`cute.copy` collectively from a converged warp. The compiler - elects one issuing lane for this operation; do not wrap the call in - :func:`cute.arch.elect_one`. An outer election would leave only one lane - able to reach the compiler-generated full-warp election, creating an - invalid synchronization that can deadlock. With NVVM diagnostics enabled, - the compiler rejects this pattern. + The compiler does not perform thread election for this operation. See the `PTX documentation `__. """ @@ -1850,12 +1940,7 @@ class CopyBulkS2SOp(CopyOp): """ Bulk copy asynchronous SMEM CTA to Cluster Copy Operation. - Invoke :func:`cute.copy` collectively from a converged warp. The compiler - elects one issuing lane for this operation; do not wrap the call in - :func:`cute.arch.elect_one`. An outer election would leave only one lane - able to reach the compiler-generated full-warp election, creating an - invalid synchronization that can deadlock. With NVVM diagnostics enabled, - the compiler rejects this pattern. + The compiler does not perform thread election for this operation. See the `PTX documentation `__. """ diff --git a/python/CuTeDSL/cutlass/cute/nvgpu/cpasync/helpers.py b/python/CuTeDSL/cutlass/cute/nvgpu/cpasync/helpers.py index 106ecef375..ca27ed5e9c 100644 --- a/python/CuTeDSL/cutlass/cute/nvgpu/cpasync/helpers.py +++ b/python/CuTeDSL/cutlass/cute/nvgpu/cpasync/helpers.py @@ -47,8 +47,14 @@ CopyBulkTensorTileS2GOp, CopyBulkTensorIm2ColS2GOp, CopyReduceBulkTensorTileS2GOp, + CopyBulkTensor2DScatter4S2GOp, + CopyBulkTensor2DGather4G2SOp, + CopyBulkTensor2DGather4G2SMulticastOp, CopyBulkTensorTileG2SNonExecTrait, CopyBulkTensorTileG2SMulticastNonExecTrait, + CopyBulkTensor2DGather4G2SNonExecTrait, + CopyBulkTensor2DGather4G2SMulticastNonExecTrait, + CopyBulkTensor2DScatter4S2GNonExecTrait, CopyBulkTensorTileS2GNonExecTrait, CopyReduceBulkTensorTileS2GNonExecTrait, CopyBulkTensorIm2ColG2SNonExecTrait, @@ -166,6 +172,9 @@ def __new_from_mlir_values__(self, values: List[Any]) -> "TmaInfo": CopyBulkTensorIm2ColG2SOp, CopyBulkTensorIm2ColS2GOp, CopyReduceBulkTensorTileS2GOp, + CopyBulkTensor2DScatter4S2GOp, + CopyBulkTensor2DGather4G2SOp, + CopyBulkTensor2DGather4G2SMulticastOp, ] @@ -423,6 +432,7 @@ def make_tiled_tma_atom( cta_tiler: Tiler, num_multicast: int = 1, *, + gmem_coord_tensor: Optional[Tensor] = None, internal_type: Optional[Type[Numeric]] = None, loc: Optional[ir.Location] = None, ip: Optional[ir.InsertionPoint] = None, @@ -463,6 +473,8 @@ def make_tiled_tma_atom( :type cta_tiler: Tiler :param num_multicast: The multicast factor :type num_multicast: int + :param gmem_coord_tensor: The GMEM index tensor for gather4/scatter4 mode (required for gather4 and scatter4 ops); see :func:`tma_partition` for layout conventions and restrictions + :type gmem_coord_tensor: Optional[Tensor] :param internal_type: Optional internal data type to use when the tensor data type is not supported by the TMA unit :type internal_type: Type[Numeric] :return: A TmaInfo containing the Copy Atom, TMA tensor, and SMEM layout @@ -560,6 +572,23 @@ def make_tiled_tma_atom( res[1], stored_smem_layout, ) + elif isinstance(op, CopyBulkTensor2DScatter4S2GOp): + if gmem_coord_tensor is None: + raise ValueError("gmem_coord_tensor is required for scatter4 TMA ops") + res = _cute_nvgpu_ir.atom_make_non_exec_2d_scatter4_tma_store( + cast(Any, gmem_tensor).value, + gmem_coord_tensor.layout, + smem_layout, + cta_v_map, + tma_format=tma_format, + loc=loc, + ip=ip, + ) + return TmaInfo( + atom.CopyAtom(op, CopyBulkTensor2DScatter4S2GNonExecTrait(res[0])), + res[1], + stored_smem_layout, + ) elif isinstance(op, CopyBulkTensorTileS2GOp): res = _cute_nvgpu_ir.atom_make_non_exec_tiled_tma_store( cast(Any, gmem_tensor).value, @@ -589,6 +618,54 @@ def make_tiled_tma_atom( res[1], stored_smem_layout, ) + elif isinstance(op, CopyBulkTensor2DGather4G2SOp): + if gmem_coord_tensor is None: + raise ValueError("gmem_coord_tensor is required for gather4 TMA ops") + if num_multicast != 1: + raise ValueError( + f"expects num_multicast to be 1 for non multicast gather4 G2S copies, " + f"but got {num_multicast}" + ) + res = _cute_nvgpu_ir.atom_make_non_exec_2d_gather4_tma_load( + gmem_tensor.value, + gmem_coord_tensor.layout, + smem_layout, + cta_v_map, + op._to_ir(), + num_multicast=num_multicast, + tma_format=tma_format, + loc=loc, + ip=ip, + ) + return TmaInfo( + atom.CopyAtom(op, CopyBulkTensor2DGather4G2SNonExecTrait(res[0])), + res[1], + stored_smem_layout, + ) + elif isinstance(op, CopyBulkTensor2DGather4G2SMulticastOp): + if gmem_coord_tensor is None: + raise ValueError("gmem_coord_tensor is required for gather4 TMA ops") + if num_multicast < 1: + raise ValueError( + f"expects num_multicast to be >= 1 for multicast gather4 G2S copies, " + f"but got {num_multicast}" + ) + res = _cute_nvgpu_ir.atom_make_non_exec_2d_gather4_tma_load( + gmem_tensor.value, + gmem_coord_tensor.layout, + smem_layout, + cta_v_map, + op._to_ir(), + num_multicast=num_multicast, + tma_format=tma_format, + loc=loc, + ip=ip, + ) + return TmaInfo( + atom.CopyAtom(op, CopyBulkTensor2DGather4G2SMulticastNonExecTrait(res[0])), + res[1], + stored_smem_layout, + ) else: raise ValueError(f"expects a bulk tensor (TMA) Copy Op, but got {op}") diff --git a/python/CuTeDSL/cutlass/cute/nvgpu/helpers.py b/python/CuTeDSL/cutlass/cute/nvgpu/helpers.py index 121fbb38d7..5b2188b476 100644 --- a/python/CuTeDSL/cutlass/cute/nvgpu/helpers.py +++ b/python/CuTeDSL/cutlass/cute/nvgpu/helpers.py @@ -32,6 +32,10 @@ CopyBulkTensorTileG2SNonExecTrait, CopyBulkTensorTileG2SMulticastOp, CopyBulkTensorTileG2SMulticastNonExecTrait, + CopyBulkTensorTileS2GOp, + CopyBulkTensorTileS2GNonExecTrait, + CopyReduceBulkTensorTileS2GOp, + CopyReduceBulkTensorTileS2GNonExecTrait, CopyBulkTensorIm2ColG2SOp, CopyBulkTensorIm2ColG2SNonExecTrait, CopyBulkTensorIm2ColG2SMulticastOp, @@ -43,6 +47,7 @@ __all__ = [ "make_tiled_tma_atom_A", "make_tiled_tma_atom_B", + "make_tiled_tma_atom_C", "make_im2col_tma_atom_A", ] @@ -348,6 +353,135 @@ def make_tiled_tma_atom_B( stored_smem_layout, ) +@dsl_user_op +def make_tiled_tma_atom_C( + op: Union[ + CopyReduceBulkTensorTileS2GOp, + CopyBulkTensorTileS2GOp, + ], + gmem_tensor: Tensor, + smem_layout: Union[Layout, ComposedLayout], + mma_tiler_mnk: Shape, + tiled_mma: atom.TiledMma, + *, + internal_type: Optional[Type[Numeric]] = None, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, +) -> TmaInfo: + """ + Makes a TMA Copy atom mapping to ``.tile`` mode for ``cp.async.bulk.tensor`` PTX operation + accounting for the MN projections of the TiledMMA for C tensor store/reduce operations. + + Given + + - a GMEM tensor + - a SMEM layout + - a MMA Tiler + - a TiledMma + - a Cluster-level shape + + this function figures out the bulk tensor asynchronous store/reduce instruction to use with the + maximum "TMA vector length" to copy tiles from the SMEM buffer to the GMEM tensor with the + provided layout and consistent with the provided Tiler & tiled_mma (considering the M-mode & N-mode). + + This function returns two results: + + 1. the Copy Atom + 2. the so-called TMA tensor used to map logical coordinates of the GMEM tensor to coordinates + that the TMA unit can consume. TMA tensors have so-called basis stride elements so that the + associated layout can output coordinates. Otherwise, TMA tensors can be partitioned + similarly to any other CuTe tensors using the algebra. + + :param op: The Copy Operation to construct an Atom for + :type op: Union[CopyReduceBulkTensorTileS2GOp, CopyBulkTensorTileS2GOp] + :param gmem_tensor: The GMEM tensor to be loaded by this copy atom + :type gmem_tensor: Tensor + :param smem_layout: Shared memory layout to load the tensor into (PDSL) + :type smem_layout: Union[Layout, ComposedLayout] + :param mma_tiler_mnk: The MMA Tiler shape (TILE_M, TILE_N, TILE_K) in MNK dimensions + :type mma_tiler_mnk: Shape + :param tiled_mma: The TiledMMA that will consume the load as operands + :type tiled_mma: atom.TiledMma + :param internal_type: Optional element-format override used when the + tensor element type does not match the copy type + :type internal_type: Type[Numeric] + :return: A TmaInfo containing the Copy Atom, TMA tensor, and SMEM layout + :rtype: TmaInfo + + """ + smem_rank = core.rank(smem_layout) + assert smem_rank == 3 or smem_rank == 4, ( + "a_smem_layout must be non-staged (atom, rest_m, rest_k) " + "or staged (atom, rest_m, rest_k, stage), " + f"but got rank = {smem_rank}" + ) + + # Keep the original SMEM layout object for later retrieval at Python level. + stored_smem_layout = smem_layout + + # Slice the smem_layout if it is staged + if smem_rank == 4: + smem_layout = core.select(smem_layout, mode=[0, 1, 2]) + + ident = core.make_identity_layout(gmem_tensor.shape, loc=loc, ip=ip) + mma_mnk: Any = mma_tiler_mnk + mma_tiler_mn = mma_mnk[:2] + g_tile = core.composition(ident, mma_tiler_mn, loc=loc, ip=ip) + cta_v_map: Any = tiled_mma._thrfrg_C(g_tile) + cta_v_map = core.get(cta_v_map, mode=[1]) + cta_v_map = core.dice(cta_v_map, (1, (1,) * core.rank(g_tile))) + + smem_for_ir: Any = smem_layout + if isinstance(smem_for_ir, core._ComposedLayout): + smem_for_ir = smem_for_ir.value + + tma_format = None + if internal_type is not None: + itype: Any = internal_type + if not isinstance(internal_type, NumericMeta): + raise TypeError(f"internal_type must be a Numeric, but got {internal_type}") + + gmem_tensor_element_type = gmem_tensor.element_type + assert not is_int_tuple_type(gmem_tensor_element_type) + + tma_format = _cute_nvgpu_ir.TmaDataFormat( + _cute_nvgpu_ir.get_default_tma_format(itype.mlir_type, False) + ) + + # res[0] = the IR Value for the non-executable atom instance + # res[1] = the IR Value for the associated TMA tensor + if isinstance(op, CopyReduceBulkTensorTileS2GOp): + res = _cute_nvgpu_ir.atom_make_non_exec_tiled_tma_reduce( + cast(Any, gmem_tensor).value, + smem_for_ir, + cta_v_map, + op._to_ir(), + tma_format=tma_format, + loc=loc, + ip=ip, + ) + return TmaInfo( + atom.CopyAtom(op, CopyReduceBulkTensorTileS2GNonExecTrait(res[0])), + res[1], + stored_smem_layout, + ) + + assert isinstance(op, CopyBulkTensorTileS2GOp) + res = _cute_nvgpu_ir.atom_make_non_exec_tiled_tma_store( + cast(Any, gmem_tensor).value, + smem_for_ir, + cta_v_map, + tma_format=tma_format, + loc=loc, + ip=ip, + ) + return TmaInfo( + atom.CopyAtom(op, CopyBulkTensorTileS2GNonExecTrait(res[0])), + res[1], + stored_smem_layout, + ) + + @dsl_user_op def make_im2col_tma_atom_A( diff --git a/python/CuTeDSL/cutlass/cute/runtime.py b/python/CuTeDSL/cutlass/cute/runtime.py index db2f89059d..f46e0093cf 100644 --- a/python/CuTeDSL/cutlass/cute/runtime.py +++ b/python/CuTeDSL/cutlass/cute/runtime.py @@ -931,11 +931,12 @@ def __get_mlir_types__(self) -> list[ir.Type]: # ------------------------------------------------------------------------- -# Register TensorAdapter for numpy/torch tensors lazily by type name, so -# importing this module never pays for importing numpy or torch itself. +# Register TensorAdapter for NumPy, PyTorch, and JAX tensors lazily by type name, +# so importing this module never imports the frameworks themselves. # ------------------------------------------------------------------------- JitArgAdapterRegistry.register_jit_arg_adapter("numpy.ndarray", lazy=True)( TensorAdapter ) JitArgAdapterRegistry.register_jit_arg_adapter("torch.Tensor", lazy=True)(TensorAdapter) +JitArgAdapterRegistry.register_jit_arg_adapter("jax.Array", lazy=True)(TensorAdapter) diff --git a/python/CuTeDSL/cutlass/cute/tensor.py b/python/CuTeDSL/cutlass/cute/tensor.py index b60e98941e..1e1a1bcabb 100644 --- a/python/CuTeDSL/cutlass/cute/tensor.py +++ b/python/CuTeDSL/cutlass/cute/tensor.py @@ -94,6 +94,8 @@ cvt_f32x4_to_fnv8e5m3x4, ) +from cutlass._mlir_helpers.dominance import value_reaches_current_ip + __all__ = [ "TensorSSA", @@ -159,14 +161,6 @@ class _Tensor(Tensor): # replacement. _pyir_ref_supported = True - # A ``_Tensor`` is a *descriptor* over memory (an iterator/pointer + - # layout), recomputable wherever its SSA value dominates. This flag tells - # the staged-tracking layer to rematerialize its ref at the use site inside - # a nested region rather than routing it through an entry-block poison - # (which leaks for cross-region reads). Declared here so the lower DSL - # layer stays decoupled from the cute dialect's MLIR type classes. - _pyir_memref_backed = True - @dsl_user_op def __init__( self, @@ -198,17 +192,7 @@ def __init__( raise TypeError(f"Expected ir.Value or _Tensor, got {type(value)}") # Set iterator - iter_val = _cute_ir.get_iter(self.value, loc=loc, ip=ip) - if isinstance(iter_val, Pointer): - self._iterator = iter_val - elif isinstance(iter_val.type, _cute_ir.ArithTupleIteratorType): - itup_val = _cute_ir.deref_arith_tuple_iter(iter_val) - self._iterator = _unpack_x_tuple(itup_val) # type: ignore[assignment] - elif isinstance(iter_val, ir.Value): - # SMEM descriptor iterator requires specific vec_mode layout configuration - self._iterator = iter_val - else: - raise TypeError(f"unsupported iterator type, got {type(iter_val)}") + self._iterator = self._derive_iterator(loc=loc, ip=ip) # Set dtype if self._dtype is None: @@ -222,6 +206,36 @@ def __init__( else: raise TypeError(f"unsupported iterator type, got {type(self.iterator)}") + def _derive_iterator( + self, + *, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, + ) -> Union[Pointer, IntTuple, ir.Value]: + """Emit ``cute.get_iter`` for this tensor and classify the result. + + The emitted op is *positional*: it is only usable where its block + dominates, whereas ``self.value`` may dominate a far wider scope. A + wrapper constructed inside a nested region therefore carries a + region-local iterator even when its memref is function-scope; if the + wrapper outlives the region, the cached iterator becomes unusable. + ``iterator`` re-derives through here whenever the cached one cannot + reach the use site. + + :raises TypeError: If iterator type is not supported + """ + iter_val = _cute_ir.get_iter(self.value, loc=loc, ip=ip) + if isinstance(iter_val, Pointer): + return iter_val + elif isinstance(iter_val.type, _cute_ir.ArithTupleIteratorType): + itup_val = _cute_ir.deref_arith_tuple_iter(iter_val) + return _unpack_x_tuple(itup_val) + elif isinstance(iter_val, ir.Value): + # SMEM descriptor iterator requires specific vec_mode layout configuration + return iter_val + else: + raise TypeError(f"unsupported iterator type, got {type(iter_val)}") + def __repr__(self) -> str: return self.__str__() @@ -469,6 +483,11 @@ def type(self) -> ir.Type: @property @lru_cache_ir() def iterator(self) -> Union[Pointer, IntTuple]: + if not value_reaches_current_ip(self._iterator): + # Minted in a region this use site cannot reach. ``self.value`` is + # the tensor's anchor and does dominate, so re-derive rather than + # hand back an operand that would fail the dominance verifier. + self._iterator = self._derive_iterator() return self._iterator @property diff --git a/python/CuTeDSL/cutlass/cutlass_dsl/_launch_facts_metadata.py b/python/CuTeDSL/cutlass/cutlass_dsl/_launch_facts_metadata.py new file mode 100644 index 0000000000..ab680c1bfd --- /dev/null +++ b/python/CuTeDSL/cutlass/cutlass_dsl/_launch_facts_metadata.py @@ -0,0 +1,40 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: LicenseRef-NvidiaProprietary +# +# Use of this software is governed by the terms and conditions of the +# NVIDIA End User License Agreement (EULA), available at: +# https://docs.nvidia.com/cutlass/latest/media/docs/pythonDSL/license.html +# +# Any use, reproduction, disclosure, or distribution of this software +# and related documentation outside the scope permitted by the EULA +# is strictly prohibited. + +"""Shared IR metadata contract for private CuTe launch-fact providers. + +The launch backend and provider-facing CuTe API both consume this attribute. +Keeping its wire names and schema version here prevents the producer and +consumer from drifting without introducing a ``cutlass_dsl``/``cutlass.cute`` +import cycle. +""" + +from cutlass._mlir import ir + + +# CUDA launch dimensions occupy unsigned 32-bit fields. +CUDA_LAUNCH_DIM_MAX = (1 << 32) - 1 + +LAUNCH_FACTS_SCHEMA_VERSION = 1 +LAUNCH_FACTS_SCHEMA_VERSION_FIELD = "schema_version" + +LAUNCH_FACTS_ATTR = "cutlass_launch_facts" +EXACT_BLOCK_DIM_FIELD = "exact_block_dim" +EXACT_GRID_DIM_FIELD = "exact_grid_dim" +EXACT_CLUSTER_DIM_FIELD = "exact_cluster_dim" +COOPERATIVE_LAUNCH_FIELD = "cooperative_launch" +CLUSTER_LAUNCH_FIELD = "cluster_launch" + + +def int64_attr(value: int) -> ir.IntegerAttr: + """Build the signed-width-neutral integer attribute used by the contract.""" + + return ir.IntegerAttr.get(ir.IntegerType.get_signless(64), value) diff --git a/python/CuTeDSL/cutlass/cutlass_dsl/cuda_jit_executor.py b/python/CuTeDSL/cutlass/cutlass_dsl/cuda_jit_executor.py index 0bb32488aa..60bdd7f26f 100644 --- a/python/CuTeDSL/cutlass/cutlass_dsl/cuda_jit_executor.py +++ b/python/CuTeDSL/cutlass/cutlass_dsl/cuda_jit_executor.py @@ -321,7 +321,11 @@ def to(self, device: Optional[int] = None) -> JitExecutor: cuda_library, ) - return JitExecutor(self.jit_module, None, self.jit_time_profiling) + # F-SPEC: the derived executor validates re-entries against the + # same sealed record as the handle it came from. + return JitExecutor( + self.jit_module, None, self.jit_time_profiling, spec=self._pyir_spec + ) @property def library(self) -> "cuda_runtime.cudaLibrary_t": diff --git a/python/CuTeDSL/cutlass/cutlass_dsl/cutlass.py b/python/CuTeDSL/cutlass/cutlass_dsl/cutlass.py index 5c6ec8004f..232364345a 100644 --- a/python/CuTeDSL/cutlass/cutlass_dsl/cutlass.py +++ b/python/CuTeDSL/cutlass/cutlass_dsl/cutlass.py @@ -16,7 +16,7 @@ # Local module imports from types import GenericAlias, SimpleNamespace, UnionType -from typing_extensions import deprecated +from typing_extensions import deprecated, override from typing import ( Callable, Generator, @@ -84,7 +84,11 @@ active_env_manager, ) from ..base_dsl.diagnostics import DiagId, find_user_source_location -from ..base_dsl.env_manager import EnvironmentVarManager, get_bool_env_var +from ..base_dsl.env_manager import ( + env_var, + EnvironmentVarManager, + get_bool_env_var, +) from ..base_dsl.utils.logger import log from ..base_dsl.utils.tree_utils import ( Leaf, @@ -145,6 +149,19 @@ LoopUnroll, ) +from ._launch_facts_metadata import ( + CLUSTER_LAUNCH_FIELD, + COOPERATIVE_LAUNCH_FIELD, + CUDA_LAUNCH_DIM_MAX, + EXACT_BLOCK_DIM_FIELD, + EXACT_CLUSTER_DIM_FIELD, + EXACT_GRID_DIM_FIELD, + LAUNCH_FACTS_ATTR, + LAUNCH_FACTS_SCHEMA_VERSION, + LAUNCH_FACTS_SCHEMA_VERSION_FIELD, + int64_attr, +) + # ============================================================================= # Cutlass DSL Device Info # ============================================================================= @@ -369,6 +386,98 @@ def _build_kernel_attrs(config: BaseDSL.LaunchConfig) -> dict: return kernel_attrs +def _is_pyir_watched_meta(value: object) -> bool: + """Return whether ``value`` may be promoted from a PyIR meta value. + + Watched integers and booleans deliberately subclass ``int`` and are not + recognized by ``is_dynamic_expression``. Treating their trace-time shadow + as exact would leave stale launch metadata if PyIR later promotes the slot. + """ + + try: + from ..base_dsl.pyir_runtime import _WatchedM + except ImportError: + return False + return isinstance(value, _WatchedM) + + +def _static_launch_dim(dims: Sequence[Any]) -> tuple[int, int, int] | None: + """Return a verified static CUDA dimension, or ``None`` for SSA values.""" + + normalized: list[int] = [] + for dim in dims: + value = dim.value if isinstance(dim, Integer) else dim + if _is_pyir_watched_meta(value): + return None + if is_dynamic_expression(value): + return None + if isinstance(value, bool) or not isinstance(value, Integral): + return None + value = int(value) + if value <= 0 or value > CUDA_LAUNCH_DIM_MAX: + return None + normalized.append(value) + if len(normalized) != 3: + return None + return tuple(normalized) # type: ignore[return-value] + + +def _static_launch_bool(value: object) -> bool | None: + """Return a verified static launch Boolean, or ``None`` for SSA values.""" + + if isinstance(value, Boolean): + value = value.value + if _is_pyir_watched_meta(value): + return None + # Static DSL booleans use the integer representation inherited from + # ``Integer`` even though their constructor canonicalizes through bool. + if isinstance(value, int) and value in (0, 1): + return bool(value) + return value if isinstance(value, bool) else None + + +def _build_launch_facts_attr(config: BaseDSL.LaunchConfig) -> ir.DictAttr: + """Encode exact facts from the same ``LaunchConfig`` used for launch. + + Mixed preferred/fallback cluster launches omit ``exact_cluster_dim`` when + the runtime may select different shapes. Dynamic fields are omitted rather + than represented as exact facts. + """ + + cluster_launch = config.has_cluster or config.has_fallback_cluster + facts: dict[str, ir.Attribute] = { + LAUNCH_FACTS_SCHEMA_VERSION_FIELD: int64_attr(LAUNCH_FACTS_SCHEMA_VERSION), + CLUSTER_LAUNCH_FIELD: ir.BoolAttr.get(cluster_launch), + } + cooperative_launch = _static_launch_bool(config.cooperative) + if cooperative_launch is not None: + facts[COOPERATIVE_LAUNCH_FIELD] = ir.BoolAttr.get(cooperative_launch) + block = _static_launch_dim(config.block) + if block is not None: + facts[EXACT_BLOCK_DIM_FIELD] = ir.DenseI64ArrayAttr.get(block) + grid = _static_launch_dim(config.grid) + if grid is not None: + facts[EXACT_GRID_DIM_FIELD] = ir.DenseI64ArrayAttr.get(grid) + + exact_cluster: tuple[int, int, int] | None = None + if config.has_cluster: + assert config.cluster is not None + if not config.has_fallback_cluster: + exact_cluster = _static_launch_dim(config.cluster) + else: + assert config.fallback_cluster is not None + preferred = _static_launch_dim(config.cluster) + fallback = _static_launch_dim(config.fallback_cluster) + if preferred is not None and preferred == fallback: + exact_cluster = preferred + elif config.has_fallback_cluster: + assert config.fallback_cluster is not None + exact_cluster = _static_launch_dim(config.fallback_cluster) + if exact_cluster is not None: + facts[EXACT_CLUSTER_DIM_FIELD] = ir.DenseI64ArrayAttr.get(exact_cluster) + return ir.DictAttr.get(facts) + + class CutlassBaseDSL(BaseDSL): """This abstract class provides a DSL for Cutlass.""" @@ -380,7 +489,7 @@ class CutlassBaseDSL(BaseDSL): @staticmethod def _make_kernel_decorator( target_cls: type["CutlassBaseDSL"], - frame: Any, + location: DSLLocation, *dargs: Any, **dkwargs: Any, ) -> Any: @@ -391,13 +500,13 @@ def _make_kernel_decorator( ``CuteExperimentalDSL.kernel``: when ``attributes`` is supplied, the resulting decorator stamps the spec onto the function via ``target_cls._KERNEL_ATTR_SPEC_FIELD`` before running the normal - jit wrapping logic. The caller is responsible for capturing the - user's source frame so source locations point to the call site - rather than this helper. + jit wrapping logic. The caller is responsible for resolving the + user's call site, so source locations point there rather than at + this helper. """ attr_spec = dkwargs.pop("attributes", None) kernel_decorator = BaseDSL.jit_runner( - target_cls, "_kernel_helper", frame, *dargs, **dkwargs + target_cls, "_kernel_helper", location, *dargs, **dkwargs ) if attr_spec is None: return kernel_decorator @@ -525,7 +634,12 @@ def _collect_extra_kernel_value_attrs( raw_attrs = self._collect_raw_kernel_attrs_from_decorator(func_body, func_args) if not raw_attrs: return {} + return self._convert_extra_kernel_value_attrs(raw_attrs) + def _convert_extra_kernel_value_attrs( + self, raw_attrs: dict[str, Any] + ) -> dict[str, ir.Attribute]: + """Validate and convert resolved kernel attributes to MLIR attributes.""" converted_attrs: dict[str, ir.Attribute] = {} for key, value in raw_attrs.items(): if key not in self._ALLOWED_EXTRA_KERNEL_VALUE_ATTRS: @@ -639,7 +753,7 @@ def _generate_kernel_attrs(self, config: BaseDSL.LaunchConfig) -> dict: f"Expect LaunchConfig for @kernel, but got {type(config)}" ) - ret = {} + ret = {LAUNCH_FACTS_ATTR: _build_launch_facts_attr(config)} if not config.has_max_number_threads(): block_str = ", ".join(map(str, config.block)) has_dynamic = any(is_dynamic_expression(dim) for dim in config.block) @@ -1730,11 +1844,12 @@ class CuTeDSLEnvironmentManager(EnvironmentVarManager): programs without an explicit per-compile opt-out (default: False). """ + use_extension_compiler: bool = env_var( + "USE_EXTENSION_COMPILER", affects_compile=True, default=False + ) + def __init__(self, prefix: str = "CUTE_DSL") -> None: super().__init__(prefix) - self.use_extension_compiler = get_bool_env_var( - f"{prefix}_USE_EXTENSION_COMPILER", False - ) class CuTeDSL(CutlassBaseDSL): @@ -1743,6 +1858,13 @@ class CuTeDSL(CutlassBaseDSL): """ _env_class = CuTeDSLEnvironmentManager + _ALLOWED_EXTRA_KERNEL_VALUE_ATTRS: frozenset[str] = frozenset( + { + "lir.tma_update_mode", + "lir.tma_override_mode", + } + ) + _KERNEL_ATTR_SPEC_FIELD: Optional[str] = "_cute_kernel_attributes" envar: CuTeDSLEnvironmentManager def __init__(self) -> None: @@ -1778,8 +1900,15 @@ def jit(cls, *dargs: Any, **dkwargs: Any) -> Any: # Capture the user's call site (mirroring BaseDSL.jit) so that # decorator source locations are reported relative to the caller, # not this override. - frame = inspect.currentframe().f_back # type: ignore[union-attr] - return BaseDSL.jit_runner(target_cls, "_func", frame, *dargs, **dkwargs) + return BaseDSL.jit_runner( + target_cls, + "_func", + BaseDSL.get_location_from_frame( + inspect.currentframe().f_back # type: ignore[union-attr] + ), + *dargs, + **dkwargs, + ) def _should_use_extension_compilation(self) -> bool: """Return whether the current compilation uses the extension compiler.""" @@ -1790,6 +1919,19 @@ def _should_use_extension_compilation(self) -> bool: # all explicit and architecture-wide opt-outs above win. return self.envar.use_extension_compiler + @override + def _collect_extra_kernel_value_attrs( + self, func_body: Callable[..., None], func_args: tuple, func_kwargs: dict + ) -> dict[str, ir.Attribute]: + """Collect extension kernel attributes only for extension compilation.""" + del func_kwargs + raw_attrs = self._collect_raw_kernel_attrs_from_decorator(func_body, func_args) + if not raw_attrs: + return {} + if not self._should_use_extension_compilation(): + raise DSLUserCodeError(DiagId.CONFIG_ATTRIBUTES_UNSUPPORTED) + return self._convert_extra_kernel_value_attrs(raw_attrs) + def _get_pipeline(self, pipeline: Optional[str]) -> str: if self._should_use_extension_compilation(): return self._get_extension_pipeline(pipeline) @@ -1834,10 +1976,10 @@ def kernel(cls, *dargs: Any, **dkwargs: Any) -> Any: and non-experimental decorations on the same kernel triggers preprocessor mismatches. attributes : optional - Kernel-level attribute spec, only supported when routing to - ``CuteExperimentalDSL`` (i.e. ``is_experimental=True``). - Stamped onto the wrapped function via - ``CuteExperimentalDSL._KERNEL_ATTR_SPEC_FIELD``. + Kernel-level attribute spec. Ordinary ``CuTeDSL`` kernels support + the extension attribute allowlist when extension compilation is selected; + ``CuteExperimentalDSL`` kernels continue to use their existing + always-on extension route. """ is_experimental = dkwargs.pop("is_experimental", False) # CuteExperimentalDSL is defined later in this module; the @@ -1848,11 +1990,13 @@ def kernel(cls, *dargs: Any, **dkwargs: Any) -> Any: # Capture the user's call site here rather than relying on a # nested method, so source locations point at the caller rather # than this override. - current_frame = inspect.currentframe() - assert current_frame is not None - frame = current_frame.f_back return CutlassBaseDSL._make_kernel_decorator( - target_cls, frame, *dargs, **dkwargs + target_cls, + BaseDSL.get_location_from_frame( + inspect.currentframe().f_back # type: ignore[union-attr] + ), + *dargs, + **dkwargs, ) def generate_func_op( @@ -2070,6 +2214,15 @@ def _device_func_impl( self, funcBody: Callable[..., Any], *args: Any, **kwargs: Any ) -> Any: if ir.Context.current is not None and ir.InsertionPoint.current is not None: + # PyIR trace-arg intake for the INLINE device-function call. + # Only a missing pyir layer is tolerable; a real intake error + # must surface, not be swallowed. + try: + from ..base_dsl.pyir_runtime import _pyir_register_trace_args + except ImportError: + pass + else: + _pyir_register_trace_args(args, kwargs) return funcBody(*args, **kwargs) # Device functions have no host entry point — never create a JIT engine. @@ -2199,6 +2352,17 @@ def _instantiate_bare(arg: Any) -> Any: from ..base_dsl.multi_stage_manager import isolated_region with isolated_region(): + # PyIR trace-arg intake: register the device-function trace's block-arg reconstructions. + # Only a missing pyir layer is tolerable; a + # real intake error must surface. + try: + from ..base_dsl.pyir_runtime import ( + _pyir_register_trace_args, + ) + except ImportError: + pass + else: + _pyir_register_trace_args(ir_args, ir_kwargs) result = funcBody(*ir_args, **ir_kwargs) # Emit the cuda.return terminator (shared with @@ -2337,10 +2501,14 @@ def kernel(cls, *dargs: Any, **dkwargs: Any) -> Any: # which would record *this* frame instead of the user's source # location (f_back would land in this override rather than in # the user file). - current_frame = inspect.currentframe() - assert current_frame is not None - frame = current_frame.f_back - return CutlassBaseDSL._make_kernel_decorator(cls, frame, *dargs, **dkwargs) + return CutlassBaseDSL._make_kernel_decorator( + cls, + BaseDSL.get_location_from_frame( + inspect.currentframe().f_back # type: ignore[union-attr] + ), + *dargs, + **dkwargs, + ) def _generate_kernel_attrs(self, config: BaseDSL.LaunchConfig) -> dict: import re diff --git a/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_ast_decorators.py b/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_ast_decorators.py index 4d2a9da8cb..846f44e03a 100644 --- a/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_ast_decorators.py +++ b/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_ast_decorators.py @@ -26,6 +26,11 @@ from ..base_dsl.dsl import is_dynamic_expression from .._mlir_helpers.arith import ArithValue from ..base_dsl.ast_helpers import * # noqa: F401,F403 +from ..base_dsl.pyir_runtime import ( + PYIR_REGION_ATTR_WRITES_ATTR, + PYIR_REGION_FREE_CALLS_ATTR, + PYIR_REGION_METHOD_CALLS_ATTR, +) from ..base_dsl.utils.logger import log from ..base_dsl import typing as t from ..base_dsl.typing import Boolean, Numeric, as_numeric, _binary_op_type_promote @@ -81,8 +86,19 @@ class ScfGenerator: Encapsulates common scf dialect functionality: pack, unpack, and SCF execution. """ - def __init__(self) -> None: - pass + # The (base, attr) pairs a staged loop body's ORIGINAL code assigns directly + # (preprocessor-tagged declared fact); set by the loop executor before building + region_attr_writes: "Optional[tuple[tuple[str, str], ...]]" = None + # The (base, method) pairs a staged loop body's ORIGINAL code calls on + # Name-rooted receiver paths (preprocessor-tagged; completed via class facts). + region_method_calls: "Optional[tuple[tuple[str, str], ...]]" = None + # The (func_name, arg_base) pairs a staged loop body's ORIGINAL code calls + # as bare-Name free functions with a Name-rooted first argument + # (preprocessor-tagged; completed via the callee's first-parameter facts). + region_free_calls: "Optional[tuple[tuple[str, str], ...]]" = None + # The staged loop-body function itself: a read-only captured receiver + # name resolves through its closure at region entry. + region_body_func: "Optional[Callable[..., Any]]" = None @staticmethod def _normalize_region_result_to_list(region_result: Any) -> List[Any]: @@ -275,10 +291,73 @@ def scf_execute_dynamic( log().debug("Completed scf.%s \n[%s]", op_type_name, op) - # 4) Pack final results + # 4) Pack final results. + # + # Rebind pass-through loop-carried slots to their INIT value instead of + # the op result: a slot every region terminator forwards unchanged (the + # terminator operand IS the entry block argument) has a result that is + # value-equal to its init for every trip count, and the init is the + # only form guaranteed to dominate the enclosing scope. The loop may + # sit inside a region-opening context manager (``with elect_one():``) + # that the AST-level value threading cannot see; a result-backed rebind + # leaks this region's SSA outward, and the enclosing staged-CF executor + # then yields it from a non-dominating region. + # Operand/argument accessors return value-caster wrappers (ArithValue, + # cute pointers) whose ``==`` stages an IR comparison; SSA identity must + # go through the base ir.Value equality on the unwrapped values. + def _same_ssa(a: object, b: object) -> bool: + def unwrap(v: object) -> Optional[ir.Value]: + if isinstance(v, ir.Value): + return v + inner = getattr(v, "value", None) + return inner if isinstance(inner, ir.Value) else None + + raw_a, raw_b = unwrap(a), unwrap(b) + if raw_a is None or raw_b is None: + return False + return ir.Value.__eq__(ir.Value(raw_a), ir.Value(raw_b)) is True + + # Accessing operands/arguments runs the registered value casters, and + # wrapper construction (e.g. the memref _Tensor caster) emits IR at the + # current insertion point -- park the insertion point at the region's + # terminator during inspection so those emissions land as valid (dead) + # in-region ops instead of parent-block ops referencing region values. + def _terminator_operands_and_args( + block: ir.Block, skip_args: int = 0 + ) -> "tuple[List[Any], List[Any]]": + term = block.operations[len(block.operations) - 1] + with ir.InsertionPoint(term): + return list(term.operands), list(block.arguments)[skip_args:] + + results: List[ir.Value] = list(op.results) + if len(results) == len(ir_values): + if op_type_name == "for": + # Block args are [induction var, carried...]. + yield_operands, carried_args = _terminator_operands_and_args( + op.regions[0].blocks[0], skip_args=1 + ) + for i, (yielded, entry) in enumerate( + zip(yield_operands, carried_args) + ): + if _same_ssa(yielded, entry): + results[i] = ir_values[i] + elif op_type_name == "while": + cond_operands, before_args = _terminator_operands_and_args( + op.regions[0].blocks[0] + ) + yield_operands, after_args = _terminator_operands_and_args( + op.regions[1].blocks[0] + ) + for i, (yielded, entry) in enumerate(zip(yield_operands, after_args)): + # cond_operands[0] is the loop condition; forwards start at 1. + if _same_ssa(yielded, entry) and _same_ssa( + cond_operands[i + 1], before_args[i] + ): + results[i] = ir_values[i] + assert isinstance(pytree_def, PyTreeDef) final_results = cutlass_dsl.pack_from_irvalue( - op.results, pytree_def, mix_iter_args, full_write_args_count + results, pytree_def, mix_iter_args, full_write_args_count ) # 5) Return in a nice pattern @@ -323,6 +402,12 @@ def _loop_execute_range_dynamic( Example: build an scf.for with optional unroll, using our universal helper. """ scf_gen = _create_control_flow_generator() + # The (base, attr) pairs the loop body's ORIGINAL code assigns directly (declared + # syntactic fact tagged by the PyIR preprocessor); the PyIR generator adopts each + scf_gen.region_attr_writes = getattr(func, PYIR_REGION_ATTR_WRITES_ATTR, None) + scf_gen.region_method_calls = getattr(func, PYIR_REGION_METHOD_CALLS_ATTR, None) + scf_gen.region_free_calls = getattr(func, PYIR_REGION_FREE_CALLS_ATTR, None) + scf_gen.region_body_func = func def create_for_op(dyn_yield_ops: List[ir.Value]) -> ir.Operation: for d in dyn_yield_ops: @@ -491,7 +576,6 @@ def _if_execute_dynamic( mix_yield_args: List[object] = [], full_write_args_count: int = 0, mix_yield_arg_names: List[str] = [], - if_constexpr: Optional[bool] = None, ) -> object: """ Build an scf.if with optional else, using our universal helper. @@ -627,20 +711,16 @@ def before_block_builder( op: ir.Operation, block_args: List[ir.Value], _: List[ir.Value], - pytree_def: Optional[PyTreeDef], + pytree_def: PyTreeDef, mix_iter_args: List[Any], full_write_args_count: int, ) -> Any: # Build the before (condition) block - if pytree_def is None: - # PyIR mode: pass original objects directly - flat_args = list(mix_iter_args) - else: - flat_args = list( - cutlass_dsl.pack_from_irvalue( - block_args, pytree_def, mix_iter_args, full_write_args_count - ) + flat_args = list( + cutlass_dsl.pack_from_irvalue( + block_args, pytree_def, mix_iter_args, full_write_args_count ) + ) log().debug("before block args: %s", flat_args) @@ -680,20 +760,16 @@ def after_block_builder( op: ir.Operation, block_args: List[ir.Value], _: List[ir.Value], - pytree_def: Optional[PyTreeDef], + pytree_def: PyTreeDef, mix_iter_args: List[object], full_write_args_count: int, ) -> object: # Build the after (body) block - if pytree_def is None: - # PyIR mode: pass original objects directly - flat_args = list(mix_iter_args) - else: - flat_args = list( - cutlass_dsl.pack_from_irvalue( - block_args, pytree_def, mix_iter_args, full_write_args_count - ) + flat_args = list( + cutlass_dsl.pack_from_irvalue( + block_args, pytree_def, mix_iter_args, full_write_args_count ) + ) log().debug("after block args: %s", flat_args) diff --git a/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_ast_decorators_pyir.py b/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_ast_decorators_pyir.py index 5890bbba05..c2ac904edb 100644 --- a/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_ast_decorators_pyir.py +++ b/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_ast_decorators_pyir.py @@ -19,20 +19,100 @@ """ import builtins +import types from typing import Any, Callable, Dict, List, Optional from cutlass._mlir import ir from cutlass._mlir.dialects import scf +from ..base_dsl.common import ( + DSLUserCodeError, + get_current_env_manager, + is_auto_m2s_enabled, +) +from ..base_dsl.diagnostics import DiagId from ..base_dsl.multi_stage_manager import enter_staged_cf, exit_staged_cf from ..base_dsl.pyir_runtime import ( + CF_ATTR_FIRST_DEF, + PYIR_REGION_ATTR_WRITES_ATTR, + PYIR_REGION_FREE_CALLS_ATTR, + PYIR_REGION_METHOD_CALLS_ATTR, + _PyirScopeGuard, + _implements_dynamic_expression, + _instance_storage_items, + _is_staged_value, + _is_untracked_post_region_binding, + _load_as_dsl, + _make_slot_key, + _meta_promote_slot, + _op_has_enclosing_loop, _pyir_auto_load_arg, + _pyir_boundary_bind_storage, + _pyir_bump_region_epoch, + _pyir_region_entry_pop, + _pyir_region_entry_push, + _pyir_gather_captured_leaf_holders, _pyir_lookup_slot_from_value, + _pyir_loop_carried_meta_needs_rebind, + _PYIR_LAST_NOTIN_COMPARE, + _PYIR_CF_ATTR_FIRST_DEFS, + _PYIR_FOLD_FIRSTDEF_STACK, + _PYIR_GATE_ONCE_INIT_SLOTS, + _slot_first_def_inside_cf, + _pyir_pop_loop_body_scope, + _pyir_push_loop_body_scope, + _pyir_rebind_local_place_cells_post_region, + _pyir_rebind_carried_leaves_post_loop, + _pyir_rebind_unstored_scalar_carries_post_loop, + _pyir_reload_stale_staged_attr_leaves, + _pyir_repair_captured_escaped_leaves, + _pyir_repair_region_escaped_leaves, + _pyir_repair_region_escaped_tuple_leaves, + _pyir_setattr_raw, + _pyir_restore_meta_numeric_leaves, + _pyir_restore_region_meta, + _pyir_snapshot_meta_numeric_leaves, + _pyir_snapshot_region_arg, + _pyir_snapshot_region_meta, + _pyir_snapshot_registry_meta_slots, + _pyir_store_back_unstored_slot_carries_sweep, + _pyir_carry_if_region_mutated_leaves, + _pyir_carry_loop_body_mutated_leaves, + _pyir_unwrap_meta_primitive, + _pyir_verify_opaque_owner_audits, + _pyir_verify_registry_meta_slots, + _raw_backing_ir_value, + _same_ir_value, + _slot_refs, + _stage_meta_compound_leaves, +) +from ..base_dsl.pyir_call_boundary import _pyir_boundary_module_is_user +from ..base_dsl.pyir_class_facts import ( + register_class_fact_modules, + transitive_write_facts, ) from ..base_dsl.typing import as_numeric from ..base_dsl.utils.logger import log from .cutlass_ast_decorators import ScfGenerator +# Sentinel distinguishing "constant-``if`` fold did not fire" from a real folded +# return of ``None`` (a folded arm with no write_args legitimately returns ``None``). +_NO_FOLD = object() + +# ``arith.cmpi`` predicate mnemonics the constant folder models. +_CMPI_PREDICATES = frozenset( + ("eq", "ne", "slt", "sle", "sgt", "sge", "ult", "ule", "ugt", "uge") +) + + +# Register modules whose holder classes mutate self-fields from jit bodies +# nested in plain methods, which decoration-time intake alone would miss. +register_class_fact_modules( + "cutlass.cute.experimental", + "cutlass.pipeline", + "cutlass.utils", +) + class PyIRScfGenerator(ScfGenerator): """ @@ -48,6 +128,65 @@ class PyIRScfGenerator(ScfGenerator): into proper SCF operations with iter_args. """ + def __init__(self) -> None: + super().__init__() + # Publish the per-DSL constructor namespace for the M2S deprecation + # message; other DSLs never build this generator and keep the default. + env_manager = get_current_env_manager() + if env_manager is not None: + env_manager.dsl_constructor_namespace = "cute" + + def _resolve_declared_base( + self, + base: str, + mix_iter_args: "List[object]", + mix_iter_arg_names: "List[str]", + ) -> Any: + """The object the dotted receiver path *base* names at region entry: + the root resolves as a loop write arg, else the staged body function's + closure binding (a read-only captured receiver is a free variable of + the body function), else its global binding; each further segment + resolves through declared instance storage only (never getattr, so a + property getter cannot execute here).""" + root, _, rest = base.partition(".") + obj = self._resolve_declared_root(root, mix_iter_args, mix_iter_arg_names) + for seg in rest.split(".") if rest else (): + if obj is None: + return None + storage = _instance_storage_items(obj) + if storage is None or seg not in storage: + return None + obj = storage[seg] + return obj + + def _resolve_declared_root( + self, + root: str, + mix_iter_args: "List[object]", + mix_iter_arg_names: "List[str]", + ) -> Any: + """The object the bare name *root* binds at region entry: a loop write + arg, else the body function's closure cell, else its global binding (a + free variable never falls through to a same-named global).""" + try: + bi = mix_iter_arg_names.index(root) + except ValueError: + bi = -1 + if 0 <= bi < len(mix_iter_args): + return mix_iter_args[bi] + fn = getattr(self, "region_body_func", None) + code = getattr(fn, "__code__", None) + if code is not None: + free = code.co_freevars + cells = getattr(fn, "__closure__", None) or () + if root in free and len(cells) == len(free): + try: + return cells[free.index(root)].cell_contents + except ValueError: + return None + fglobals = getattr(fn, "__globals__", None) + return fglobals.get(root) if fglobals is not None else None + def scf_execute_dynamic( self, op_type_name: str, @@ -88,6 +227,248 @@ def scf_execute_dynamic( block_term_op_builder, ) + @staticmethod + def _int_from_const_attr(attr: Any) -> Optional[int]: + """Extract the integer value of an ``arith.constant``'s ``value`` attribute, or + ``None`` for a non-integer attribute.""" + # The binding exposes no ``isinstance`` on these attr classes; probe by name + # then construct (guarded, since the wrong class raises): ``i1`` -> BoolAttr. + tyname = type(attr).__name__ + if tyname == "BoolAttr": + try: + return 1 if ir.BoolAttr(attr).value else 0 + except Exception: + return None + try: + return int(ir.IntegerAttr(attr).value) + except Exception: + return None + + @staticmethod + def _integer_value_bit_width(val: "ir.Value") -> Optional[int]: + """Return the bit width of an integer-typed ``ir.Value``, or ``None`` when the + type is not a plain ``IntegerType`` (or is unavailable).""" + try: + ty = val.type + if ir.IntegerType.isinstance(ty): + return ir.IntegerType(ty).width + except Exception: + return None + return None + + @staticmethod + def _reinterpret_at_value_type(v: int, val: "ir.Value") -> Optional[int]: + """Signed reinterpretation of *v* at *val*'s own integer width (the + invariant every fold result maintains); ``i1`` normalizes to 0/1.""" + width = PyIRScfGenerator._integer_value_bit_width(val) + if width is None: + return None + v &= (1 << width) - 1 + if width > 1 and v >= 1 << (width - 1): + v -= 1 << width + return v + + @staticmethod + def _eval_constant_int(val: "ir.Value", _depth: int = 0) -> Optional[int]: + """Constant-fold an integer ``ir.Value`` defined by a small tree of ``arith`` ops + over compile-time constants; every result is the signed reinterpretation + at its own type width (i1 as 0/1); return its int, else None.""" + if _depth > 64: + return None + wrap = PyIRScfGenerator._reinterpret_at_value_type + try: + owner = val.owner + op = getattr(owner, "operation", owner) + if op is None or not hasattr(op, "name"): + return None + name = op.name + if name == "arith.constant": + c = PyIRScfGenerator._int_from_const_attr(op.attributes["value"]) + return None if c is None else wrap(c, val) + operands = list(op.operands) + + def ev(i: int) -> Optional[int]: + return PyIRScfGenerator._eval_constant_int(operands[i], _depth + 1) + + if name == "arith.trunci": + a = ev(0) + return None if a is None else wrap(a, val) + if name == "arith.extsi": + # Sign-extension preserves the signed value (the invariant). + return ev(0) + if name == "arith.extui": + # Zero-extension takes the operand's UNSIGNED reinterpretation. + a = ev(0) + if a is None: + return None + opw = PyIRScfGenerator._integer_value_bit_width(operands[0]) + return None if opw is None else a & ((1 << opw) - 1) + if name == "arith.andi": + a, b = ev(0), ev(1) + return None if a is None or b is None else wrap(a & b, val) + if name == "arith.ori": + a, b = ev(0), ev(1) + return None if a is None or b is None else wrap(a | b, val) + if name == "arith.xori": + a, b = ev(0), ev(1) + return None if a is None or b is None else wrap(a ^ b, val) + if name == "arith.addi": + a, b = ev(0), ev(1) + return None if a is None or b is None else wrap(a + b, val) + if name == "arith.subi": + a, b = ev(0), ev(1) + return None if a is None or b is None else wrap(a - b, val) + if name == "arith.muli": + a, b = ev(0), ev(1) + return None if a is None or b is None else wrap(a * b, val) + if name == "arith.cmpi": + a, b = ev(0), ev(1) + if a is None or b is None: + return None + # Recover the comparison predicate mnemonic. + pstr = str(op.attributes["predicate"]).strip().rstrip(">") + mnem = pstr.split()[-1] if pstr else "" + # Conservative decline: an unrecognized mnemonic leaves the + # ``if`` as runtime control flow. + if mnem not in _CMPI_PREDICATES: + return None + # Unsigned predicates compare the operands' UNSIGNED + # interpretation (a recovered int can be negative). + ua, ub = a, b + if mnem in ("ult", "ule", "ugt", "uge"): + width = PyIRScfGenerator._integer_value_bit_width(operands[0]) + if width is not None: + mask = (1 << width) - 1 + ua, ub = a & mask, b & mask + table = { + "eq": a == b, + "ne": a != b, + "slt": a < b, + "sle": a <= b, + "sgt": a > b, + "sge": a >= b, + "ult": ua < ub, + "ule": ua <= ub, + "ugt": ua > ub, + "uge": ua >= ub, + } + if mnem not in table: + return None + return 1 if table[mnem] else 0 + if name == "arith.select": + c = ev(0) + if c is None: + return None + return ev(1) if c else ev(2) + return None + except Exception: + return None + + @staticmethod + def _constant_i1_value(cond: "ir.Value") -> Optional[bool]: + """Return the Python bool of a statically-decidable ``i1`` *cond*, or ``None`` + when *cond* is not a compile-time constant.""" + v = PyIRScfGenerator._eval_constant_int(cond) + return None if v is None else bool(v) + + def _maybe_fold_constant_if( + self, + op_type_name: str, + op: "ir.Operation", + region_builders: List[Callable[..., Any]], + mix_iter_args: List[object], + ) -> Any: + """Fold a constant-condition ``scf.if`` to in-line execution.""" + if op_type_name != "if": + return _NO_FOLD + op_view = getattr(op, "operation", op) + try: + if len(op_view.operands) == 0: + return _NO_FOLD + cond = op_view.operands[0] + except Exception: + return _NO_FOLD + const = self._constant_i1_value(cond) + if const is None: + return _NO_FOLD + # Gate-once evidence (see pyir_state): the compare choke's one-slot + # register describes THIS fold's predicate when it was a single compare. + _gate_cand = _PYIR_LAST_NOTIN_COMPARE[0] + _PYIR_LAST_NOTIN_COMPARE[0] = None + if _gate_cand is not None and _gate_cand[2] is not bool(const): + _gate_cand = None # register predates this fold's predicate + + # then == region_builders[0]; else (if present) == region_builders[1]. + live_builder = None + if const and len(region_builders) >= 1: + live_builder = region_builders[0] + elif (not const) and len(region_builders) >= 2: + live_builder = region_builders[1] + + # Erase the empty scf.if shell BEFORE tracing the live arm so the arm + # body is emitted at the outer scope (not after a dangling op). + try: + op_view.erase() + except Exception: + return _NO_FOLD + + log().info( + "[pyir] constant-if folded (cond=%s); running %s arm in-line", + const, + "then" if const else "else", + ) + + if live_builder is None: + # Constant-false ``if`` with no ``else``: nothing executes. + if not mix_iter_args: + return + if len(mix_iter_args) == 1: + return mix_iter_args[0] + return mix_iter_args + + # Trace the live arm in place: ``pytree_def=None`` selects the PyIR builder + # branch (shared objects), and the op handle is ``None`` (shell was erased). + _collect_gate_once = bool(const) and _gate_cand is not None + if _collect_gate_once: + _PYIR_FOLD_FIRSTDEF_STACK.append([]) + try: + region_result = live_builder(None, [], [], None, mix_iter_args, 0) + finally: + if _collect_gate_once: + _fold_firstdefs = _PYIR_FOLD_FIRSTDEF_STACK.pop() + try: + _gset, _gkey, _ = _gate_cand + if _gkey in _gset: # the arm latched the gate + for _fd in _fold_firstdefs: + if ( + isinstance(_fd, tuple) + and len(_fd) == 2 + and _fd[0] == CF_ATTR_FIRST_DEF + ): + # Latched gate = once-per-key init: discharge + # the read-before-set record. + _PYIR_CF_ATTR_FIRST_DEFS.pop(_fd[1], None) + continue + _slot_first_def_inside_cf[_fd] = False + _PYIR_GATE_ONCE_INIT_SLOTS.add(_fd) + except Exception: + pass + if region_result is not None: + result_list = ( + region_result + if isinstance(region_result, (list, tuple)) + else [region_result] + ) + for idx, val in enumerate(result_list): + if idx < len(mix_iter_args) and val is not None: + mix_iter_args[idx] = val + + if not mix_iter_args: + return + if len(mix_iter_args) == 1: + return mix_iter_args[0] + return mix_iter_args + def _scf_execute_pyir( self, op_type_name: str, @@ -112,8 +493,23 @@ def _scf_execute_pyir( # Create SCF op with zero iter_args — pyir refs handle all state op = create_op_func([]) + # Meta-advance witness for LOOP bodies: snapshot the registered holders' + # meta attrs so the close can refuse an un-instrumented advance. + _meta_slot_snapshot = ( + _pyir_snapshot_registry_meta_slots() + if op_type_name in ("for", "while") + else None + ) log().debug("Generated scf.%s (pyir) \n[%s]", op_type_name, op) + # Constant-condition ``if`` is META CF: run the single reachable arm IN-LINE and + # erase the shell, else a staged ``scf.if %true`` traps a leaf mutation (miscompile). + folded = self._maybe_fold_constant_if( + op_type_name, op, region_builders, mix_iter_args + ) + if folded is not _NO_FOLD: + return folded + # For multi-region ops (scf.if), save originals so the else-block # doesn't see MLIR values from inside the then-block's region. is_multi_region = op_type_name == "if" and len(region_builders) > 1 @@ -121,38 +517,439 @@ def _scf_execute_pyir( list(mix_iter_args) if is_multi_region else None ) + # Snapshot each iter_arg's value-tree leaves at the OUTER IP; bodies run on + # shared objects, so the post-region repair reverts any region-trapped escapee. + leaf_snapshots: List[Optional[List["ir.Value"]]] = [ + _pyir_snapshot_region_arg(arg) for arg in mix_iter_args + ] + + # Snapshot each iter_arg's meta-scalar leaves at the OUTER scope so + # ``_pyir_restore_region_meta`` reverts an in-place meta mutation that escapes. + meta_snapshots: List[Optional[List[tuple]]] = [ + _pyir_snapshot_region_meta(arg) for arg in mix_iter_args + ] + + # Also snapshot value-tree leaves of CAPTURED read-only objects (not in + # ``mix_iter_args``); a shared sub-object leaf mutation else traps an SSA a sibling reads. + captured_leaf_ids = {id(arg) for arg in mix_iter_args} + captured_leaf_snapshot: Optional[List[tuple]] = ( + _pyir_gather_captured_leaf_holders(exclude_ids=captured_leaf_ids) + ) + + # Gate the Numeric-leaf region-escape revert to ``if`` only: in a loop a Numeric + # leaf is a genuine carried iter_arg, so reverting it would drop the carry. + allow_numeric_revert = op_type_name == "if" + + # Snapshot WRITE-ONLY literal-origin Numeric (uncarried counter) leaves so each + # sibling re-bakes its own constant; ``if`` only (a loop legitimately carries it). + meta_numeric_snapshots: List[Optional[List[tuple]]] = ( + [_pyir_snapshot_meta_numeric_leaves(arg) for arg in mix_iter_args] + if allow_numeric_revert + else [None] * len(mix_iter_args) + ) + + # For/while: carry an in-place-mutated raw leaf through a ``pyir.ref`` so it lifts + # to an iter_arg (``while`` opaque-only, to not double-carry M2S scalar state). + carries_loop_leaves = op_type_name in ("for", "while") + loop_leaf_opaque_only = op_type_name == "while" + loop_leaf_records: List[tuple] = [] + # An un-instrumented self-attr scalar mutation (``self.attr += step`` inside a + # plain, non-jit method): stored back into the existing slot ref the body loads + unstored_scalar_carry_records: List[tuple] = [] + + # Region index of the loop BODY (for: 0; while after/body: 1; non-loop: -1). + # Bounds the per-iteration-reset recording scope to the body only. + loop_body_region_index = ( + 0 if op_type_name == "for" else (1 if op_type_name == "while" else -1) + ) + + # Region-entry ledger adoption from the body's DECLARED write facts: + # direct ``(base, attr)`` assign pairs, plus ``(base, method)`` call + # pairs completed into receiver-attr writes through the callee's + # class facts (a method-call mutation is syntactically invisible to + # the direct-assign collector). + declared_attr_write_records: List[tuple] = [] + # Unpromoted-write audit rows (holder, attr, entry value, label): + # recorded at region entry, checked at region close. Covers nested + # direct writes, while-body direct writes, and jit-callee receiver + # writes the call boundary never pre-stages (default mode only). + deep_attr_write_audits: List[tuple] = [] + if op_type_name in ("for", "while"): + + def _adopt_declared_attr_write( + _holder: Any, + _attr: str, + _plain_meta_only: bool = False, + _stage_compound_only: bool = False, + ) -> None: + """Get-or-mint the ledger cell for ``(holder, attr)`` at the + OUTER IP so the body's stores carry it as a loop carry.""" + # A value-protocol holder carries its state through its own + # extract/reconstruct pair: the callee-fact call sites must + # not stage (and so clone) its fields behind that protocol. + if _stage_compound_only and _implements_dynamic_expression( + _holder + ): + return + # The declared pair names a ledger place only when it is an + # INSTANCE-STORAGE slot of the holder. + _storage = _instance_storage_items(_holder) + if _storage is None or _attr not in _storage: + return + _cur = _storage[_attr] + # Method-completed pairs mutate in UN-CHOKED code, so nothing + # stores into a cell minted over an already-STAGED value (the + # staged value's own carry machinery owns it); a direct-assign + # pair stores back through the assign choke, which serves + # both cells, so staged current values stay admitted there. + if _plain_meta_only and _is_staged_value(_cur): + return + _py = _pyir_unwrap_meta_primitive(_cur) + if _py is None or type(_py) not in (bool, int, float): + # Compound field (tuple/list/dict/object): stage each leaf + # so the per-leaf carry can thread the body's rebuild. + # AUTO_M2S only; the default mode refuses in pyir_assign. + # Not for the non-user callee domain: its writes reach no + # choke and no boundary replay, so staged leaves would + # never see the rebuild. + if not _plain_meta_only and is_auto_m2s_enabled(): + _staged = _stage_meta_compound_leaves(_cur) + if _staged is not None: + _pyir_boundary_bind_storage(_holder, _attr, _staged) + return + if _stage_compound_only: + return # scalar: this call site never mints scalar cells + try: + _place = _make_slot_key(None, _holder, _attr) + except Exception: + _place = None + if _place is None or _slot_refs.get(_place) is not None: + return # get-or-mint anti-twin: an existing cell stands + _ref = _meta_promote_slot( + _place, + _py, + promoted_value=_cur if _cur is not _py else None, + display_name=f"{type(_holder).__name__}.{_attr}", + ) + if _ref is None: + return + declared_attr_write_records.append((_holder, _attr, _ref, _place)) + + _audit_seen: "set[tuple[int, str]]" = set() + + def _arm_unpromoted_write_audit( + _holder: Any, _attr: str, _label: str + ) -> None: + """Record (holder, attr, entry value, label) for the close-time + unpromoted-write audit; plain meta primitive leaves only (a + staged leaf's own carry machinery owns it).""" + if _holder is None: + return + _key = (id(_holder), _attr) + if _key in _audit_seen: + return + _storage = _instance_storage_items(_holder) + if _storage is None or _attr not in _storage: + return + _cur = _storage[_attr] + if _is_staged_value(_cur): + return + _py = _pyir_unwrap_meta_primitive(_cur) + if _py is None or type(_py) not in (bool, int, float): + return + _audit_seen.add(_key) + deep_attr_write_audits.append((_holder, _attr, _py, _label)) + + def _callee_write_attrs(_callee: Any) -> "frozenset[str]": + try: + _writes, _ = transitive_write_facts(_callee) + except Exception: + return frozenset() + return _writes + + facts = getattr(self, "region_attr_writes", None) or () + for _base, _attr in facts: + if "." in _base or op_type_name == "while": + # Nested direct write (`c.sub.n = ...`), or any direct + # write in a while body (which has no adoption path): + # record the entry value and re-check it at region close; + # refusing here would break a leaf a staged consumption + # promotes mid-body, which is carried fine. Under + # AUTO_M2S the read-side promotion rescues the shape. + if is_auto_m2s_enabled(): + if "." not in _base: + # A COMPOUND field rebuilt in a while body has no + # read-side rescue: stage its leaves so the choke + # decomposition carries the rebuild, as in for. + _holder = self._resolve_declared_base( + _base, mix_iter_args, mix_iter_arg_names + ) + if _holder is not None and not isinstance( + _holder, (bool, int, float, str) + ): + _adopt_declared_attr_write( + _holder, _attr, _stage_compound_only=True + ) + continue + _holder = self._resolve_declared_base( + _base, mix_iter_args, mix_iter_arg_names + ) + _arm_unpromoted_write_audit(_holder, _attr, f"{_base}.{_attr}") + continue + try: + _bi = mix_iter_arg_names.index(_base) + except ValueError: + continue + if _bi >= len(mix_iter_args): + continue + _holder = mix_iter_args[_bi] + if _holder is None or isinstance(_holder, (bool, int, float, str)): + continue + _adopt_declared_attr_write(_holder, _attr) + + method_facts = getattr(self, "region_method_calls", None) or () + for _base, _method in method_facts: + _holder = self._resolve_declared_base( + _base, mix_iter_args, mix_iter_arg_names + ) + if _holder is None or isinstance( + _holder, + (bool, int, float, str, bytes, type, types.ModuleType), + ): + continue + # Resolve the named def through the class MRO -- no instance + # getattr, so a property getter cannot execute here; a + # staticmethod/classmethod first parameter is not the receiver. + _target = None + for _klass in type(_holder).__mro__: + _cand = _klass.__dict__.get(_method) + if _cand is not None: + _target = _cand + break + if _target is None or isinstance( + _target, (staticmethod, classmethod, property) + ): + continue + _target = getattr(_target, "__func__", _target) + if not callable(_target) or not hasattr(_target, "__code__"): + continue + # USER-module defs are rewritten by the preprocessor, so their + # receiver writes reach the chokes directly; the declared + # completion imports write knowledge only for defs the + # rewrite never sees (the boundary's non-user callee domain). + if _pyir_boundary_module_is_user(getattr(_target, "__module__", None)): + # A jit-decorated callee traces inline, skipping the call + # boundary's imminent-write pre-stage, so in the default + # mode its choked receiver writes run on the plain Python + # value and bake: audit its write set at region close. + # (A plain callee's boundary pre-stage mints the cell, so + # the audit passes it untouched.) + if is_auto_m2s_enabled(): + # A COMPOUND field the callee rebuilds reaches the + # choke (jit: directly; plain: boundary replay), but + # only staged leaves make the choke carry instead of + # refuse: stage them at entry, as for direct writes. + for _attr in _callee_write_attrs( + _target.__get__(_holder, type(_holder)) + ): + _adopt_declared_attr_write( + _holder, _attr, _stage_compound_only=True + ) + else: + for _attr in _callee_write_attrs( + _target.__get__(_holder, type(_holder)) + ): + _arm_unpromoted_write_audit( + _holder, _attr, f"{_base}.{_attr}" + ) + continue + if op_type_name == "while": + # No adoption for while regions (facts newly reach them; + # adoption is a for-only behavior): audit-only. + if not is_auto_m2s_enabled(): + for _attr in _callee_write_attrs( + _target.__get__(_holder, type(_holder)) + ): + _arm_unpromoted_write_audit( + _holder, _attr, f"{_base}.{_attr}" + ) + continue + try: + _writes, _ = transitive_write_facts( + _target.__get__(_holder, type(_holder)) + ) + except Exception: + continue + for _attr in _writes: + _adopt_declared_attr_write(_holder, _attr, _plain_meta_only=True) + + free_facts = getattr(self, "region_free_calls", None) or () + for _fname, _base in free_facts: + # Free call with a Name-rooted first argument (`step(c)`): + # a jit-decorated callee traces inline (no boundary + # pre-stage), so its first-parameter writes bake in the + # default mode; audit them at region close. Under AUTO_M2S + # a USER callee's compound rebuild reaches the choke (jit: + # directly; plain: boundary replay): stage its leaves at + # entry so the choke decomposition carries it. + _holder = self._resolve_declared_base( + _base, mix_iter_args, mix_iter_arg_names + ) + if _holder is None or isinstance( + _holder, + (bool, int, float, str, bytes, type, types.ModuleType), + ): + continue + _fn = self._resolve_declared_root( + _fname, mix_iter_args, mix_iter_arg_names + ) + _fn = getattr(_fn, "__func__", _fn) + if ( + _fn is None + or isinstance(_fn, type) + or not callable(_fn) + or not hasattr(_fn, "__code__") + ): + continue + if is_auto_m2s_enabled(): + if _pyir_boundary_module_is_user( + getattr(_fn, "__module__", None) + ): + for _attr in _callee_write_attrs(_fn): + _adopt_declared_attr_write( + _holder, _attr, _stage_compound_only=True + ) + continue + for _attr in _callee_write_attrs(_fn): + _arm_unpromoted_write_audit( + _holder, _attr, f"{_base}.{_attr}" + ) + + # If analogue: carry an opaque leaf advanced+trapped in an ``scf.if`` + # forward (yield, not revert) so a post-``if`` reader sees it; aliases + # excluded below. + carries_if_leaves = op_type_name == "if" + if_leaf_records: List[tuple] = [] + if_leaf_ref_cache: dict = {} + # STANDALONE ``scf.if`` only: leaf SSAs SHARED across aliases keep their revert + # (forward-carry one strands the rest); ``None`` in a loop, which carries them. + if_aliased_leaves: Optional[list] = None + if carries_if_leaves and not _op_has_enclosing_loop(op): + if_aliased_leaves = [] + if captured_leaf_snapshot: + for _ch, _ca, _cv in captured_leaf_snapshot: + _craw = _raw_backing_ir_value(_cv) + if _craw is not None: + if_aliased_leaves.append(_craw) + # Opaque leaf SSAs held by 2+ mix_iter_args objects (aliased sub-object): + # collect per-slot, then keep those seen in more than one iter_arg slot. + _per_arg_leaves: List[List[ir.Value]] = [] + for _li in builtins.range(len(leaf_snapshots)): + _ls = leaf_snapshots[_li] + _slot_leaves: List[ir.Value] = [] + if _ls: + for _h, _a, _v in _ls: + _raw = _raw_backing_ir_value(_v) + if isinstance(_raw, ir.Value): + _slot_leaves.append(_raw) + _per_arg_leaves.append(_slot_leaves) + for _li in builtins.range(len(_per_arg_leaves)): + for _v in _per_arg_leaves[_li]: + _in_other = any( + any(_same_ir_value(_w, _v) for _w in _per_arg_leaves[_lj]) + for _lj in builtins.range(len(_per_arg_leaves)) + if _lj != _li + ) + if _in_other: + if_aliased_leaves.append(_v) + enter_staged_cf() try: for i, builder in enumerate(region_builders): - # Reset to original values before building non-first regions - # (e.g. else-block) so they don't see then-block's region- - # local MLIR values that would violate SSA dominance. + # Reconcile escaped CAPTURED leaves at EVERY region's entry (not just the + # ``i > 0`` reset) so a chained-branch shared leaf can't trap an SSA the next reads. + _pyir_repair_captured_escaped_leaves( + captured_leaf_snapshot, + f"{op_type_name} captured (region {i} entry)", + allow_numeric_revert, + ) + + # Reset to original values before non-first regions (else-block) so they + # don't see then-block region-local MLIR values (SSA dominance violation). if is_multi_region and i > 0: assert original_mix_iter_args is not None for idx in builtins.range(len(mix_iter_args)): mix_iter_args[idx] = original_mix_iter_args[idx] + # The SHARED object's leaf still carries the prior branch's trapped + # SSA, so revert any leaf not dominating the outer IP (snapshot does). + _pyir_repair_region_escaped_leaves( + mix_iter_args[idx], + leaf_snapshots[idx], + f"{op_type_name} iter_arg #{idx} (pre-region {i})", + allow_numeric_revert, + ) + # Restore meta-scalar leaves the then-block mutated so the else-block + # sees op-entry meta (both branches start equal, as in non-PyIR). + _pyir_restore_region_meta( + mix_iter_args[idx], + meta_snapshots[idx], + f"{op_type_name} iter_arg #{idx} (pre-region {i})", + ) + # Restore a write-only literal-origin Numeric counter the prior + # branch baked in so this sibling re-bakes its own (non-carried only). + _pyir_restore_meta_numeric_leaves( + meta_numeric_snapshots[idx], + f"{op_type_name} iter_arg #{idx} (pre-region {i})", + ) + # Captured objects were already reconciled by the region-entry repair + # above; the idempotent revert means a second repair here is a no-op. region = op.regions[i] block = region.blocks[0] with ir.InsertionPoint(block): block_args = list(block.arguments) - # Execute body -- AST-inserted pyir_assign/pyir_read - # handle ref creation, load, store for all mutable values. - region_result = builder( - op, - block_args, - [], # ir_values: empty (no iter_args) - None, # pytree_def: not used (pyir handles everything) - mix_iter_args, - 0, # full_write_args_count: 0 + # Bind each declared-attr-write adoption to a fresh region- + # entry load of its cell (pre-write in-body reads see it). + for _dh, _da, _dref, _dplace in declared_attr_write_records: + try: + _pyir_boundary_bind_storage( + _dh, _da, _load_as_dsl(_dref, place=_dplace) + ) + except Exception: + pass + + # Re-load any tracked object-attribute slot whose staged value + # escaped a sibling region, so an un-instrumented read in THIS + _pyir_reload_stale_staged_attr_leaves( + f"{op_type_name} region {i} entry" ) - # Update mix_iter_args from body result so that - # slot-backed objects are accessible after the loop - # (slot-backed values carry ``_mutable_ref`` / - # ``_pyir_load_version`` tags the post-loop bridge - # below re-loads from). + # Declare the loop BODY block as the innermost open body scope + # (F-BIRTHPOS): first-defs traced inside it are region-born. + is_loop_body_region = i == loop_body_region_index + if is_loop_body_region: + _pyir_push_loop_body_scope(block) + else: + # Non-body region entry (if arm / while cond) is a region-epoch + # boundary; body regions bump inside the body-scope push. + _pyir_bump_region_epoch() + _pyir_region_entry_push() + + # Execute body -- AST-inserted pyir_assign/pyir_read handle ref + # creation, load, store for all mutable values. + with _PyirScopeGuard(kind="region", inherit=True): + region_result = builder( + op, + block_args, + [], # ir_values: empty (no iter_args) + None, # pytree_def: not used (pyir handles everything) + mix_iter_args, + 0, # full_write_args_count: 0 + ) + + # Update mix_iter_args from body result so slot-backed objects (and their + # ``_mutable_ref`` / load-version tags) survive to the post-loop bridge. if region_result is not None: result_list = ( region_result @@ -163,6 +960,99 @@ def _scf_execute_pyir( if idx < len(mix_iter_args) and val is not None: mix_iter_args[idx] = val + # Post-body (still open): carry any in-place-mutated leaf through a + # ``pyir.ref`` so it lifts to an iter_arg instead of being reverted. + if carries_loop_leaves: + for idx in builtins.range(len(mix_iter_args)): + loop_leaf_records.extend( + _pyir_carry_loop_body_mutated_leaves( + mix_iter_args[idx], + leaf_snapshots[idx], + op, + block, + f"{op_type_name} iter_arg #{idx}", + opaque_only=loop_leaf_opaque_only, + # The slot to re-point post-loop for a SELF-LEAF (a whole-object + # bare ``ir.Value`` is immutable, so its carry re-binds the slot). + arg_index=idx, + # The slot's bare-name fact: a SELF-LEAF resolves its + # LOCAL place so every region adopts ONE cell. + arg_name=( + mix_iter_arg_names[idx] + if idx < len(mix_iter_arg_names) + else None + ), + ) + ) + # Same opaque-leaf carry for CAPTURED free-variable + # objects advanced in a nested loop inside this one. + _arg_covered_leaf_keys = { + (id(_h), _a) + for _ls in leaf_snapshots + if _ls + for (_h, _a, _v) in _ls + } + captured_carry_snapshot = ( + [ + rec + for rec in captured_leaf_snapshot + if (id(rec[0]), rec[1]) not in _arg_covered_leaf_keys + ] + if captured_leaf_snapshot + else None + ) + loop_leaf_records.extend( + _pyir_carry_loop_body_mutated_leaves( + None, + captured_carry_snapshot, + op, + block, + f"{op_type_name} captured", + opaque_only=loop_leaf_opaque_only, + ) + ) + # Store back an un-instrumented self-attr scalar mutation (a + # plain-method ``self.attr += step`` or state advance): the body + unstored_scalar_carry_records.extend( + _pyir_store_back_unstored_slot_carries_sweep( + op, + block, + f"{op_type_name} registry-sweep (region {i})", + ) + ) + + # Post-body (still open): close the loop-body scope declaration. + if is_loop_body_region: + _pyir_pop_loop_body_scope() + else: + # Non-body region close is a region-epoch boundary; body + # regions bump inside the body-scope pop. + _pyir_bump_region_epoch() + _pyir_region_entry_pop() + + # Post-region (still open): carry this branch's mutated leaf + # forward (yield) so the enclosing loop carries it, instead + # of reverting it. + if carries_if_leaves: + for idx in builtins.range(len(mix_iter_args)): + if_leaf_records.extend( + _pyir_carry_if_region_mutated_leaves( + mix_iter_args[idx], + leaf_snapshots[idx], + op, + block, + if_leaf_ref_cache, + f"{op_type_name} region #{i} iter_arg #{idx}", + arg_index=idx, + exclude_aliased_leaves=if_aliased_leaves, + arg_name=( + mix_iter_arg_names[idx] + if idx < len(mix_iter_arg_names) + else None + ), + ) + ) + # Terminator if builder in block_term_op_builder: block_term_op_builder[builder](region_result, 0) @@ -173,30 +1063,162 @@ def _scf_execute_pyir( log().debug("Completed scf.%s (pyir) \n[%s]", op_type_name, op) - # Emit pyir.load for slot-backed values so that returned - # values are valid at the outer scope (not trapped inside the - # scf.if/for/while body where they were defined). - # Primary slot detection routes through _pyir_lookup_slot_from_value - # (the _pyir_load_version + _mutable_ref tags on the value); when the - # inner-scope value has no load tag yet, we additionally consult the - # ``_mutable_ref`` attribute as a fallback so "slot-backed but never loaded" cases - # still get the post-loop load. The ``_mutable_ref`` marker is - # then bridged onto the loaded value so the outer scope's - # ``_pyir_auto_load_arg`` -- which consults ``_mutable_ref`` per - # the current pyir_runtime contract -- can chain loads at each - # nested-CF scope exit. + # Rebind any loop-carried leaf to a dominating post-loop ``pyir.load`` so post- + # loop reads observe the loop result (the repair below then skips it). + if carries_loop_leaves and loop_leaf_records: + rebuilt_by_index = _pyir_rebind_carried_leaves_post_loop(loop_leaf_records) + for ai, obj in rebuilt_by_index.items(): + if 0 <= ai < len(mix_iter_args): + mix_iter_args[ai] = obj + + # Rebind each stored-back scalar carry to a dominating post-loop + # ``pyir.load`` so later reads observe the carried value. + if carries_loop_leaves and unstored_scalar_carry_records: + _pyir_rebind_unstored_scalar_carries_post_loop( + unstored_scalar_carry_records + ) + + # Rebind each declared-attr-write adoption to a dominating post-loop + # load, so a post-loop reader observes the carried result. + for _dh, _da, _dref, _dplace in declared_attr_write_records: + try: + _pyir_boundary_bind_storage( + _dh, _da, _load_as_dsl(_dref, place=_dplace) + ) + except Exception: + pass + + # Rebind any if-forward-carried leaf to a dominating post-``if`` load; a WHOLE-OBJECT + # rebind re-points the slot to the object the empty-else reset had discarded. + if carries_if_leaves and if_leaf_records: + rebuilt_by_index = _pyir_rebind_carried_leaves_post_loop(if_leaf_records) + for ai, obj in rebuilt_by_index.items(): + if 0 <= ai < len(mix_iter_args): + mix_iter_args[ai] = obj + + # Post-region rebind FROM the ledger cell (clause A, read half): a bare + # LOCAL re-sources its binding from a dominating post-region load. + _pyir_rebind_local_place_cells_post_region( + mix_iter_args, mix_iter_arg_names, op + ) + + # Revert escaped leaves of CAPTURED objects to their pre-op value; ordered AFTER + # the loop-leaf rebind so a now-dominating carried leaf is not reverted. + _pyir_repair_captured_escaped_leaves( + captured_leaf_snapshot, f"{op_type_name} captured", allow_numeric_revert + ) + # After this scf op closes, re-load any tracked object-attribute slot whose + # staged value was bound inside the just-closed region (now non-dominating at + _pyir_reload_stale_staged_attr_leaves(f"{op_type_name} post-region") + + # Emit pyir.load for slot-backed values so returns are valid at the outer scope; + # bridge the ``_mutable_ref`` marker onto the load so loads chain at CF exits. for idx in builtins.range(len(mix_iter_args)): arg = mix_iter_args[idx] + # Repair value-tree leaves a region body mutated to an SSA from the now-closed + # region, which would else fail ``module.verify()`` (operand does not dominate). + arg = _pyir_repair_region_escaped_leaves( + arg, + leaf_snapshots[idx], + f"{op_type_name} iter_arg #{idx}", + allow_numeric_revert, + ) + # Re-bind a TUPLE-contained scalar value-tree leaf the body trapped: the + # value-tree walk above never collects a tuple-nested scalar, so a sibling + if if_aliased_leaves is not None: + _pyir_repair_region_escaped_tuple_leaves( + arg, op, f"{op_type_name} iter_arg #{idx}" + ) + # Restore meta-scalar leaves a region body mutated to the op-entry value, so a + # later op sharing the object observes entry meta (the non-PyIR post-op state). + arg = _pyir_restore_region_meta( + arg, meta_snapshots[idx], f"{op_type_name} iter_arg #{idx}" + ) + # Restore a write-only literal-origin Numeric counter the closed region trapped, + # so a later op sharing the object re-bakes its own constant. + _pyir_restore_meta_numeric_leaves( + meta_numeric_snapshots[idx], f"{op_type_name} iter_arg #{idx}" + ) + mix_iter_args[idx] = arg mv = _pyir_lookup_slot_from_value(arg) if mv is None: # Slot-backed but never loaded in this scope: fall back # to ``_mutable_ref`` so the bridge below still fires. mv = getattr(arg, "_mutable_ref", None) - loaded = _pyir_auto_load_arg(arg) + loaded = _pyir_auto_load_arg(arg, row_authoritative=True) if loaded is not arg: if mv is not None and _pyir_lookup_slot_from_value(loaded) is None: - object.__setattr__(loaded, "_mutable_ref", mv) + _pyir_setattr_raw(loaded, "_mutable_ref", mv) mix_iter_args[idx] = loaded + arg = loaded + + # A slot promoted to a ref inside the region but whose binding lost its ref link + # folds back stale; the ref is authoritative, so rebind to a dominating load. + cur = mix_iter_args[idx] + name = mix_iter_arg_names[idx] if idx < len(mix_iter_arg_names) else None + slot_key = _make_slot_key(name, None, None) if name else None + ref = _slot_refs.get(slot_key) if slot_key is not None else None + # Re-bind to a fresh dominating ``pyir.load %ref`` of the published slot + # when: (1) the binding lost its ref link (untracked first-def gate); or (2) + needs_rebind = ( + _is_untracked_post_region_binding(cur) + or (is_multi_region and ref is not None) + # A ``while``/``for`` carried META scalar whose slot the body stored + # into stays a meta wrapper post-loop (tracked, so the untracked gate + or ( + op_type_name in ("while", "for") + and ref is not None + and _pyir_loop_carried_meta_needs_rebind(cur, ref, op) + ) + ) + if needs_rebind and ref is not None: + # Row-authoritative reload: reconstruct from the row's store-time + # template; attaches a ``_mutable_ref`` for downstream re-loads. + mix_iter_args[idx] = _load_as_dsl(ref, place=slot_key) + log().info( + "[pyir] post-region rebind of promoted slot " + "'%s' to pyir.load (was stale %r)", + name, + cur, + ) + + # Refuse (loudly) any registered holder attr still advanced outside the ledger + # after every store-back / rebind above had its chance to cover it. + if _meta_slot_snapshot is not None: + _pyir_verify_registry_meta_slots(_meta_slot_snapshot) + + # Refuse (loudly) an audited write leaf still meta and moved at close: + # nothing promoted it, so its updates ran on the plain Python value + # and were baked at the trace value. + for _dh, _da, _entry, _label in deep_attr_write_audits: + _dstorage = _instance_storage_items(_dh) + if _dstorage is None or _da not in _dstorage: + continue + _dcur = _dstorage[_da] + if _is_staged_value(_dcur): + continue + _dpy = _pyir_unwrap_meta_primitive(_dcur) + if _dpy is None: + continue + try: + _dplace = _make_slot_key(None, _dh, _da) + except Exception: + _dplace = None + if _dplace is not None and _slot_refs.get(_dplace) is not None: + continue + try: + _changed = bool(_dpy != _entry) + except Exception: + _changed = True + if _changed: + raise DSLUserCodeError( + DiagId.PHASE_MUTATE_PYTHON, + var=_label, + ) + + # Refuse (loudly) any container key created through an opaque + # ``__setitem__`` owner in this region's bodies; retires the audits. + _pyir_verify_opaque_owner_audits() # Return in standard pattern if not mix_iter_args: @@ -221,6 +1243,19 @@ def create_while_op_pyir( log().debug("_while_execute_dynamic (PyIR path)") while_op_type_name = "while" + # The while AFTER block carries the body's declared write facts + # (preprocessor-tagged, same vocabulary as the for tag). + self.region_attr_writes = getattr( + while_after_block, PYIR_REGION_ATTR_WRITES_ATTR, None + ) + self.region_method_calls = getattr( + while_after_block, PYIR_REGION_METHOD_CALLS_ATTR, None + ) + self.region_free_calls = getattr( + while_after_block, PYIR_REGION_FREE_CALLS_ATTR, None + ) + self.region_body_func = while_after_block + _pyir_cond = [None] # list for nonlocal mutation in closures def create_while_op_impl(dyn_yield_ops: List[ir.Value]) -> ir.Operation: @@ -248,8 +1283,11 @@ def before_block_builder_pyir( def before_block_terminator_pyir( region_result: Any, full_write_args_count: int ) -> None: - # Emit scf.condition with empty pass-through args - ir_cond = as_numeric(_pyir_cond[0]).ir_value() + # Auto-load the condition before reading its SSA, so an un-instrumented + # getter's pre-loop value gets a ``pyir.load`` and carries as an iter_arg. + # The condition cell is the row the before-block just wrote. + cond = _pyir_auto_load_arg(_pyir_cond[0], row_authoritative=True) + ir_cond = as_numeric(cond).ir_value() scf.ConditionOp(ir_cond, []) def after_block_builder_pyir( @@ -263,10 +1301,9 @@ def after_block_builder_pyir( # Execute loop body. AST-inserted pyir_assign handles stores. flat_args = list(mix_iter_args) assert while_after_block is not None - while_after_block(*flat_args) - # Return None: pyir.store already updated refs. - # Default terminator in _scf_execute_pyir emits scf.YieldOp([]). - return None + # Return the body's final write-arg bindings so ``_scf_execute_pyir`` + # refreshes ``mix_iter_args`` before it carries loop-carried leaves: a leaf + return while_after_block(*flat_args) return self.scf_execute_dynamic( op_type_name=while_op_type_name, diff --git a/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_function.py b/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_function.py index 52315db0a6..d380c48638 100644 --- a/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_function.py +++ b/python/CuTeDSL/cutlass/cutlass_dsl/cutlass_function.py @@ -408,6 +408,7 @@ class CutlassCallJitCompiledFunction: "__llvm_ir__", "__mlir__", # JIT metadata / bookkeeping. + "seal_specialization", "function_name", "kernel_info", "execution_args", @@ -498,7 +499,13 @@ def to(self, device: Any = None) -> Any: "and is not supported by the cutlass compiler backend yet." ) - def get_aux_func(self, func_class: Any, kernel: Any) -> Any: + def get_aux_func( + self, + func_class: Any, + kernel: Any = None, + *, + required: bool = True, + ) -> Any: raise NotImplementedError( "CutlassCallJitCompiledFunction.get_aux_func() requires engine symbol " "lookup and is not supported by the cutlass compiler backend yet." diff --git a/python/CuTeDSL/cutlass/experimental/cuda/tensor_map.py b/python/CuTeDSL/cutlass/experimental/cuda/tensor_map.py index 08168cbe44..d6cb113cf0 100644 --- a/python/CuTeDSL/cutlass/experimental/cuda/tensor_map.py +++ b/python/CuTeDSL/cutlass/experimental/cuda/tensor_map.py @@ -45,6 +45,8 @@ Float6E3M2FN, Float6E2M3FN, Float4E2M1FNx2, + Float6E3M2FNx4, + Float6E2M3FNx4, ) import cutlass.cute as cute from cutlass.cute.core import ScaledBasis, depth, leading_dim @@ -63,9 +65,16 @@ def _stride_to_tma_units( Element units here mean units of ``element_type`` itself. For packed dtypes such as ``Float4E2M1FNx2`` / ``Float6E{3M2,2M3}FNx4``, one tensor element is already one packed storage unit. + + A dynamic stride is widened to 64 bits before scaling. Tensor strides + are commonly 32-bit values, and ``stride * width`` wraps once the + stride reaches 2^27 fp16 elements, which hands the encoder a negative + or truncated byte stride. """ - return stride * element_type.width // 128 + if isinstance(stride, int): + return stride * element_type.width // 128 + return Int64(stride) * element_type.width // 128 def _product(values: Sequence[Int8 | int]) -> Int32 | int: @@ -472,6 +481,11 @@ def _derive_tensormap_stride_dtype( if tma_format in {TensorMapDataFormat.B4X16, TensorMapDataFormat.B4X16_P64}: if dtype is Float4E2M1FNx2: return Float4E2M1FN + if tma_format == TensorMapDataFormat.B6X16_P32: + if dtype is Float6E3M2FNx4: + return Float6E3M2FN + if dtype is Float6E2M3FNx4: + return Float6E2M3FN return dtype @@ -659,6 +673,8 @@ def get_dsl_type_to_tensormap_type(dsl_type: Type[Numeric]) -> TensorMapDataType elif dsl_type in {Float6E3M2FN, Float6E2M3FN}: # 6-bit FP6 — same alignment story as FP4. return TensorMapDataType.f416u6_align16b + elif dsl_type in {Float6E3M2FNx4, Float6E2M3FNx4}: + return TensorMapDataType.f416u6_align16b raise ValueError(f"Unsupported type for TensorMap: {dsl_type}") diff --git a/python/CuTeDSL/cutlass/experimental/primitives/descriptors.py b/python/CuTeDSL/cutlass/experimental/primitives/descriptors.py index 8af9cb50a7..0a896e177e 100644 --- a/python/CuTeDSL/cutlass/experimental/primitives/descriptors.py +++ b/python/CuTeDSL/cutlass/experimental/primitives/descriptors.py @@ -55,8 +55,10 @@ Numeric, Uint8, Uint32, + Uint64, ) import cutlass.base_dsl.typing as _cutlass +from cutlass.base_dsl import target_version as _target_version from cutlass.base_dsl.typing import Array # Auto-converting proxy over the raw NVVM dialect bindings: lets the private @@ -118,6 +120,42 @@ def _tcgen05_mma_smem_desc( ) +@dsl_user_op +def _tcgen05_mma_smem_desc_v2( + start_address: int | Int32 | Uint32, + leading_dim_offset: int | Int32 | Uint32, + stride_dim_offset: int | Int32 | Uint32, + descriptor_version: int | Int8 | Uint8, + base_offset: int | Int8 | Uint8, + leading_dim_mode: int | Boolean, + k_segment_offset: int | Boolean, + swizzle_type: int | Int8 | Uint8, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Int64: + """Build an extended ``tcgen05.mma`` SMEM descriptor. + + Like :func:`_tcgen05_mma_smem_desc` with two extra descriptor fields, + ``descriptor_version`` and ``k_segment_offset``. Module-private; construct + descriptors through :meth:`Tcgen05SmemDesc.build`, the sole customer. + All offsets are in units of 16 bytes (the descriptor granule). + """ + return _cutlass.Int64( + _nvvm.tcgen05_mma_smem_desc_v2( + _cutlass.Int32(start_address), + _cutlass.Int32(leading_dim_offset), + _cutlass.Int32(stride_dim_offset), + _cutlass.Int8(descriptor_version), + _cutlass.Int8(base_offset), + _cutlass.Boolean(leading_dim_mode), + _cutlass.Boolean(k_segment_offset), + _cutlass.Int8(swizzle_type), + loc=loc, + ip=ip, + ) + ) + # ============================================================================= # Enums @@ -358,6 +396,12 @@ def advance_start_address( # ============================================================================= +# Toolchains whose assembler needs the split-word form of +# ``Tcgen05SmemDesc.advance_start_address`` (CUDA 12.x and 13.0). From 13.1 on +# the plain 64-bit add is assembled correctly and is kept for its lower op count. +_SPLIT_WORD_START_ADDRESS_ADVANCE: bool = _target_version(max_version="13.0") + + class Tcgen05SmemDesc(Int64, width=64): # type: ignore[call-arg] """SM100 ``tcgen05.mma`` shared-memory descriptor (64-bit). @@ -419,16 +463,26 @@ def advance_start_address( included, so the byte offset must be a multiple of 16. Passing a Python ``int`` that violates this raises :class:`ValueError`. - This lowers to the same encoded start-address addition as - ``desc + (byte_increment >> 4)``. The caller must ensure the - resulting start address still fits the 14-bit field. - """ - return self + _drop_low_bits(byte_increment, 4, "byte_increment") # type: ignore[return-value] + This computes ``desc + (byte_increment >> 4)``. On CUDA 13.1+ + toolchains it lowers to that plain 64-bit add; on CUDA 12.x / 13.0 + toolchains the addition is done on the low 32-bit word and the + untouched high word is re-packed, so the descriptor constant bits + never pass through a 64-bit add. The caller must ensure the resulting + start address still fits the 14-bit field. - # Keep ``desc + encoded_offset_16b`` on the base Int64 path. A masked - # ``__add__`` overload emits extra descriptor-field arithmetic in hot - # tcgen05 K loops. ``advance_start_address()`` keeps the byte-count API - # while lowering through the same fast encoded-add path. + """ + inc = _drop_low_bits(byte_increment, 4, "byte_increment") + if not _SPLIT_WORD_START_ADDRESS_ADVANCE: + return self + inc # type: ignore[return-value] + lo = Uint32(self) + Uint32(inc) + hi = (Uint64(self) >> 32) << 32 + return type(self)(Int64(hi | Uint64(lo))) # type: ignore[return-value] + + # Keep ``desc + encoded_offset_16b`` on the base Int64 path where the + # toolchain allows it. A masked ``__add__`` overload emits extra + # descriptor-field arithmetic in hot tcgen05 K loops; + # ``advance_start_address()`` keeps the byte-count API while lowering + # through the fast encoded-add path (or the split-word form above). @classmethod def build( @@ -436,9 +490,11 @@ def build( start_address: Array | cutlass.Pointer | Int32 | int, leading_byte_offset: int = 0, stride_byte_offset: int = 0, + version: Literal[0, 1, 2, 3] = 1, base_offset: Literal[0, 1, 2, 3, 4, 5, 6, 7] = 0, layout: int | Tcgen05SmemSwizzle = 0, leading_dim_mode: int = 0, + k_segment_offset: int = 0, ) -> "Tcgen05SmemDesc": """Build the SMEM descriptor via the NVVM descriptor intrinsic. @@ -482,6 +538,8 @@ def build( default unless you explicitly need the SM103+ absolute-address leading-dimension mode. + :param version: Descriptor version field. + :param k_segment_offset: B K_Segment start offset in K=128 (48B) block. """ # A swizzle-descriptor object (Tcgen05SmemSwizzle / cutlass.Swizzle / # TensorMapSwizzle) converts to the tcgen05 encoding via .to(); a plain @@ -510,14 +568,19 @@ def build( _drop_low_bits(stride_byte_offset, 4, "stride_byte_offset") ) + # The extended descriptor op carries descriptor_version and + # k_segment_offset (the K-segmented form) required by newer tcgen05 + # descriptor layouts. return cls( - _tcgen05_mma_smem_desc( - addr >> 4, - leading_dim_offset, - stride_dim_offset, - _cutlass.Int8(base_offset), - _cutlass.Boolean(leading_dim_mode), - swizzle_code, + _tcgen05_mma_smem_desc_v2( + start_address=addr >> 4, + leading_dim_offset=leading_dim_offset, + stride_dim_offset=stride_dim_offset, + descriptor_version=_cutlass.Int8(version), + base_offset=_cutlass.Int8(base_offset), + leading_dim_mode=_cutlass.Boolean(leading_dim_mode), + k_segment_offset=_cutlass.Boolean(k_segment_offset), + swizzle_type=swizzle_code, ) ) @@ -773,9 +836,10 @@ def build( assert sparse_flag in [0, 1], f"sparse_flag must be 0 or 1, got {sparse_flag}" assert saturate in [0, 1], f"saturate must be 0 or 1, got {saturate}" assert 0 <= c_format < 4, f"c_format must be in [0, 3], got {c_format}" - assert sparse_format in [0, 1], ( - f"sparse_format must be 0 or 1, got {sparse_format}" - ) + assert sparse_format in [ + 0, + 1, + ], f"sparse_format must be 0 or 1, got {sparse_format}" assert 0 <= a_format < 8, f"a_format must be in [0, 7], got {a_format}" assert 0 <= b_format < 8, f"b_format must be in [0, 7], got {b_format}" assert a_negate in [0, 1], f"a_negate must be 0 or 1, got {a_negate}" @@ -949,8 +1013,12 @@ def _input_format_from_dtype(dtype: type[Numeric]) -> int: return 1 if dtype is _cutlass.Float6E2M3FN: return 3 + if dtype is _cutlass.Float6E2M3FNx4: + return 3 if dtype is _cutlass.Float6E3M2FN: return 4 + if dtype is _cutlass.Float6E3M2FNx4: + return 4 if dtype is _cutlass.Float4E2M1FN: return 5 if dtype is _cutlass.Float4E2M1FNx2: diff --git a/python/CuTeDSL/cutlass/experimental/primitives/gpu_ops.py b/python/CuTeDSL/cutlass/experimental/primitives/gpu_ops.py index 9bea522e3a..9eea905fc7 100644 --- a/python/CuTeDSL/cutlass/experimental/primitives/gpu_ops.py +++ b/python/CuTeDSL/cutlass/experimental/primitives/gpu_ops.py @@ -73,6 +73,7 @@ def _get_smem_allocator() -> "cutlass.memory.SmemAllocator": dsl_obj = None caller = inspect.currentframe() frame = caller.f_back if caller is not None else None + del caller while frame: obj = frame.f_locals.get("self", None) if obj and isinstance(obj, CutlassBaseDSL): diff --git a/python/CuTeDSL/cutlass/experimental/primitives/nvvm_wrapper.py b/python/CuTeDSL/cutlass/experimental/primitives/nvvm_wrapper.py index 210df8eccf..0df40acf21 100644 --- a/python/CuTeDSL/cutlass/experimental/primitives/nvvm_wrapper.py +++ b/python/CuTeDSL/cutlass/experimental/primitives/nvvm_wrapper.py @@ -3439,6 +3439,40 @@ def _assert_tcgen05_swizzle(swizzle: int | Int8 | Uint8, instruction: str) -> No +@dsl_user_op +def add_packed_bf16x2_f32x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.add_packed_bf16x2_f32x2_f32x2``.""" + return _nvvm.add_packed_bf16x2_f32x2_f32x2( + T.vector(2, _cutlass.BFloat16.mlir_type), + src_a, + src_b, + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def add_packed_f16x2_f32x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.add_packed_f16x2_f32x2_f32x2``.""" + return _nvvm.add_packed_f16x2_f32x2_f32x2( + T.vector(2, _cutlass.Float16.mlir_type), + src_a, + src_b, + loc=loc, + ip=ip, + ) def _packed_f32x2_to_vec( @@ -3487,6 +3521,45 @@ def add_packed_f32x2( return _unpack_packed_f32x2(vec_res) if returns_tuple else vec_res +@dsl_user_op +def add_packed_f32x2_bf16x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + rnd: FPRoundingMode | None = None, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.add_packed_f32x2_bf16x2_f32x2``.""" + return _nvvm.add_packed_f32x2_bf16x2_f32x2( + T.vector(2, _cutlass.Float32.mlir_type), + src_a, + src_b, + rnd=_to_dialect(rnd, _FP_ROUNDING_MODE_TO_DIALECT), + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def add_packed_f32x2_f16x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + rnd: FPRoundingMode | None = None, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.add_packed_f32x2_f16x2_f32x2``.""" + return _nvvm.add_packed_f32x2_f16x2_f32x2( + T.vector(2, _cutlass.Float32.mlir_type), + src_a, + src_b, + rnd=_to_dialect(rnd, _FP_ROUNDING_MODE_TO_DIALECT), + loc=loc, + ip=ip, + ) + @dsl_user_op @@ -4304,8 +4377,6 @@ def cp_async_bulk_tensor_prefetch( :type l2_cache_hint: int or cutlass.Int64 or cutlass.Uint64, optional :raises ValueError: if the ``coordinates`` count is invalid for ``mode``. - The descriptor-override path is exposed separately as - :func:`cp_async_bulk_tensor_prefetch_override`. """ _assert_coords(coordinates, "cp.async.bulk.prefetch.tensor", mode=mode) if l2_cache_hint is not None: @@ -4833,6 +4904,70 @@ def cvt_packfloat_f32( ) +@dsl_user_op +def cvt_packfloat_f32_sf( + src_a: float | Float32, + src_b: float | Float32, + src_c: int | Int32 | Uint32, + sf: int | Int16 | Uint16, + to: CVTPackFloat, + *, + rnd: FPRoundingMode | None = None, + sat: SaturationModeKind | None = None, + relu: bool | None = None, + extract_hi: bool | None = None, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Int32: + """Wrapper over ``nvvm.cvt_packfloat_f32_sf``.""" + return _cutlass.Int32( + _nvvm.cvt_packfloat_f32_sf( + _cutlass.Float32(src_a), + _cutlass.Float32(src_b), + _cutlass.Int32(src_c), + _cutlass.Int16(sf), + _CVT_PACK_FLOAT_TO_DIALECT[to], + rnd=_to_dialect(rnd, _FP_ROUNDING_MODE_TO_DIALECT), + sat=_to_dialect(sat, _SATURATION_MODE_KIND_TO_DIALECT), + relu=relu, + extract_hi=extract_hi, + loc=loc, + ip=ip, + ) + ) + + +@dsl_user_op +def cvt_packfloat_sf( + src_a: int | Int32 | Uint32, + src_c: int | Int32 | Uint32, + sf: int | Int16 | Uint16, + from_: CVTPackFloat, + to: CVTPackFloat, + *, + rnd: FPRoundingMode | None = None, + sat: SaturationModeKind | None = None, + relu: bool | None = None, + extract_hi: bool | None = None, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Int32: + """Wrapper over ``nvvm.cvt_packfloat_sf``.""" + return _cutlass.Int32( + _nvvm.cvt_packfloat_sf( + _cutlass.Int32(src_a), + _cutlass.Int32(src_c), + _cutlass.Int16(sf), + _CVT_PACK_FLOAT_TO_DIALECT[from_], + _CVT_PACK_FLOAT_TO_DIALECT[to], + rnd=_to_dialect(rnd, _FP_ROUNDING_MODE_TO_DIALECT), + sat=_to_dialect(sat, _SATURATION_MODE_KIND_TO_DIALECT), + relu=relu, + extract_hi=extract_hi, + loc=loc, + ip=ip, + ) + ) @@ -5220,6 +5355,48 @@ def fma_packed_f32x2( return _unpack_packed_f32x2(vec_res) if returns_tuple else vec_res +@dsl_user_op +def fma_packed_f32x2_bf16x2_f32x2_f32x2( + src_a: Vector, + src_b: Vector, + src_c: Vector, + *, + rnd: FPRoundingMode | None = None, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.fma_packed_f32x2_bf16x2_f32x2_f32x2``.""" + return _nvvm.fma_packed_f32x2_bf16x2_f32x2_f32x2( + T.vector(2, _cutlass.Float32.mlir_type), + src_a, + src_b, + src_c, + rnd=_to_dialect(rnd, _FP_ROUNDING_MODE_TO_DIALECT), + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def fma_packed_f32x2_f16x2_f32x2_f32x2( + src_a: Vector, + src_b: Vector, + src_c: Vector, + *, + rnd: FPRoundingMode | None = None, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.fma_packed_f32x2_f16x2_f32x2_f32x2``.""" + return _nvvm.fma_packed_f32x2_f16x2_f32x2_f32x2( + T.vector(2, _cutlass.Float32.mlir_type), + src_a, + src_b, + src_c, + rnd=_to_dialect(rnd, _FP_ROUNDING_MODE_TO_DIALECT), + loc=loc, + ip=ip, + ) _VALID_LDST_MATRIX_NUM = frozenset({1, 2, 4}) @@ -6651,8 +6828,76 @@ def mul( ) +@dsl_user_op +def mul_packed_bf16x2_bf16x2_f16x2( + src_a: Vector, + src_b: Vector, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.mul_packed_bf16x2_bf16x2_f16x2``.""" + return _nvvm.mul_packed_bf16x2_bf16x2_f16x2( + T.vector(2, _cutlass.BFloat16.mlir_type), + src_a, + src_b, + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def mul_packed_bf16x2_f32x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.mul_packed_bf16x2_f32x2_f32x2``.""" + return _nvvm.mul_packed_bf16x2_f32x2_f32x2( + T.vector(2, _cutlass.BFloat16.mlir_type), + src_a, + src_b, + loc=loc, + ip=ip, + ) +@dsl_user_op +def mul_packed_f16x2_f16x2_bf16x2( + src_a: Vector, + src_b: Vector, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.mul_packed_f16x2_f16x2_bf16x2``.""" + return _nvvm.mul_packed_f16x2_f16x2_bf16x2( + T.vector(2, _cutlass.Float16.mlir_type), + src_a, + src_b, + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def mul_packed_f16x2_f32x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.mul_packed_f16x2_f32x2_f32x2``.""" + return _nvvm.mul_packed_f16x2_f32x2_f32x2( + T.vector(2, _cutlass.Float16.mlir_type), + src_a, + src_b, + loc=loc, + ip=ip, + ) @dsl_user_op @@ -7326,6 +7571,40 @@ def store_ext( ) +@dsl_user_op +def sub_packed_bf16x2_f32x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.sub_packed_bf16x2_f32x2_f32x2``.""" + return _nvvm.sub_packed_bf16x2_f32x2_f32x2( + T.vector(2, _cutlass.BFloat16.mlir_type), + src_a, + src_b, + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def sub_packed_f16x2_f32x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.sub_packed_f16x2_f32x2_f32x2``.""" + return _nvvm.sub_packed_f16x2_f32x2_f32x2( + T.vector(2, _cutlass.Float16.mlir_type), + src_a, + src_b, + loc=loc, + ip=ip, + ) @dsl_user_op @@ -7349,6 +7628,45 @@ def sub_packed_f32x2( ) +@dsl_user_op +def sub_packed_f32x2_bf16x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + rnd: FPRoundingMode | None = None, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.sub_packed_f32x2_bf16x2_f32x2``.""" + return _nvvm.sub_packed_f32x2_bf16x2_f32x2( + T.vector(2, _cutlass.Float32.mlir_type), + src_a, + src_b, + rnd=_to_dialect(rnd, _FP_ROUNDING_MODE_TO_DIALECT), + loc=loc, + ip=ip, + ) + + +@dsl_user_op +def sub_packed_f32x2_f16x2_f32x2( + src_a: Vector, + src_b: Vector, + *, + rnd: FPRoundingMode | None = None, + loc: ir.Location | None = None, + ip: ir.InsertionPoint | None = None, +) -> Vector: + """Wrapper over ``nvvm.sub_packed_f32x2_f16x2_f32x2``.""" + return _nvvm.sub_packed_f32x2_f16x2_f32x2( + T.vector(2, _cutlass.Float32.mlir_type), + src_a, + src_b, + rnd=_to_dialect(rnd, _FP_ROUNDING_MODE_TO_DIALECT), + loc=loc, + ip=ip, + ) + @dsl_user_op def tcgen05_alloc( @@ -7946,6 +8264,7 @@ def tcgen05_mma( enable_input_d: int | Boolean, *, collector_op: Tcgen05MMACollectorOp | None = None, + b_collector_op: Tcgen05MMACollectorOp | None = None, a_shift: bool | None = None, scale_input_d: int | Int64 | Uint64 | None = None, write_disable_mask: Vector | None = None, @@ -8073,6 +8392,15 @@ def tcgen05_mma( M-stack / M2 chain; that reuses B. :type collector_op: Tcgen05MMACollectorOp, optional + :param b_collector_op: Target-specific collector usage for operand B + (``.collector::b::*``). Same values as *collector_op*. Use only when + the target and MMA form support B-side collector reuse and the same + physical B operand is reused across the collector chain, for example an + M-stack / M2 chain. Public Blackwell B collector reuse uses + :func:`tcgen05_mma_ws` with ``collector_b_buffer`` and ``collector_op``; + do not infer plain non-WS B collector legality from A collector support. + :type b_collector_op: Tcgen05MMACollectorOp, optional + :param a_shift: When ``True``, emits the ``.ashift`` modifier. In the ``.ashift`` MMA pipeline, the shift is a **post-MMA** operation: the current MMA reads unshifted A and produces @@ -8187,6 +8515,9 @@ def tcgen05_mma( _cutlass.Int32(idesc), _cutlass.Boolean(enable_input_d), collector_op=_to_dialect(collector_op, _TCGEN05_MMA_COLLECTOR_OP_TO_DIALECT), + collector_op_b=_to_dialect( + b_collector_op, _TCGEN05_MMA_COLLECTOR_OP_TO_DIALECT + ), a_shift=a_shift, scale_input_d=scale_input_d, disable_output_lane=write_disable_mask, @@ -11242,7 +11573,11 @@ def cvt_f32x2_to_f4x2( __all__ = [ + "add_packed_bf16x2_f32x2_f32x2", + "add_packed_f16x2_f32x2_f32x2", "add_packed_f32x2", + "add_packed_f32x2_bf16x2_f32x2", + "add_packed_f32x2_f16x2_f32x2", "atomicrmw", "auto", "bar_warp_sync", @@ -11296,6 +11631,8 @@ def cvt_f32x2_to_f4x2( "cvt_f32x2_to_f8x2", "cvt_packfloat", "cvt_packfloat_f32", + "cvt_packfloat_f32_sf", + "cvt_packfloat_sf", "dot_accumulate_2way", "dot_accumulate_4way", "elect_sync", @@ -11312,6 +11649,8 @@ def cvt_f32x2_to_f4x2( "fence_sc_cluster", "fence_sync_restrict", "fma_packed_f32x2", + "fma_packed_f32x2_bf16x2_f32x2_f32x2", + "fma_packed_f32x2_f16x2_f32x2_f32x2", "fmin", "griddepcontrol", "inline_ptx", @@ -11346,6 +11685,10 @@ def cvt_f32x2_to_f4x2( "mov_b32", "mul", "mul_bf16x2", + "mul_packed_bf16x2_bf16x2_f16x2", + "mul_packed_bf16x2_f32x2_f32x2", + "mul_packed_f16x2_f16x2_bf16x2", + "mul_packed_f16x2_f32x2_f32x2", "mul_packed_f32x2", "nanosleep", "pmevent", @@ -11362,7 +11705,11 @@ def cvt_f32x2_to_f4x2( "st_bulk", "stmatrix", "store_ext", + "sub_packed_bf16x2_f32x2_f32x2", + "sub_packed_f16x2_f32x2_f32x2", "sub_packed_f32x2", + "sub_packed_f32x2_bf16x2_f32x2", + "sub_packed_f32x2_f16x2_f32x2", "tcgen05_alloc", "tcgen05_commit", "tcgen05_cp", diff --git a/python/CuTeDSL/cutlass/experimental/task_scheduling/resources.py b/python/CuTeDSL/cutlass/experimental/task_scheduling/resources.py index ccfbfee567..7fc35a8fa7 100644 --- a/python/CuTeDSL/cutlass/experimental/task_scheduling/resources.py +++ b/python/CuTeDSL/cutlass/experimental/task_scheduling/resources.py @@ -95,9 +95,8 @@ cast, get_origin, ) -import cutlass.utils.static_persistent_tile_scheduler as _static_persistent_tile_scheduler from cutlass.utils.static_persistent_tile_scheduler import ( - WorkTileInfo, + WorkTileInfo as _BaseWorkTileInfo, ) from cutlass.utils import ( StaticPersistentTileScheduler, @@ -120,76 +119,66 @@ from .pipeline_group import PipelineGroup -# TS loop-carries ``WorkTileInfo`` through staged persistent control flow and -# advances it in tail schedule entries. The upstream implementation stores -# ``tile_idx`` as one tuple field, which makes in-place updates fail IR -# dominance checks in nested dynamic regions. Keep the public CUTLASS type but -# scalarize its runtime representation while this module is loaded. -@cute.jit -def _work_tile_info_init_scalar( - self: WorkTileInfo, tile_idx: cute.Coord, is_valid_tile: cutlass.Boolean -) -> None: - m_idx, n_idx, l_idx = tile_idx # type: ignore[misc] - self._m_idx = cutlass.Int32(m_idx) # type: ignore[attr-defined] - self._n_idx = cutlass.Int32(n_idx) # type: ignore[attr-defined] - self._l_idx = cutlass.Int32(l_idx) # type: ignore[attr-defined] - self._is_valid_tile = cutlass.Boolean(is_valid_tile) - - -def _work_tile_info_extract_scalar(self: WorkTileInfo) -> list: - return [ - self._m_idx.ir_value(), # type: ignore[attr-defined] - self._n_idx.ir_value(), # type: ignore[attr-defined] - self._l_idx.ir_value(), # type: ignore[attr-defined] - self._is_valid_tile.ir_value(), - ] +class WorkTileInfo(_BaseWorkTileInfo): + """TS-local 3D work tile with scalarized loop-carried fields. + TS loop-carries work tiles through staged persistent control flow and + advances them in tail schedule entries. The released CUTLASS work-tile class + stores ``tile_idx`` as one tuple field, which makes in-place updates fail IR + dominance checks in nested dynamic regions. Keep that released class + untouched and scalarize only the TS-local subclass. + """ -def _work_tile_info_new_from_mlir_values_scalar( - self: WorkTileInfo, values: list -) -> WorkTileInfo: - assert len(values) == 4 - return WorkTileInfo( - ( - cutlass.Int32(values[0]), - cutlass.Int32(values[1]), - cutlass.Int32(values[2]), - ), - cutlass.Boolean(values[3]), - ) + @cute.jit + def __init__(self, tile_idx: cute.Coord, is_valid_tile: cutlass.Boolean) -> None: + m_idx, n_idx, l_idx = tile_idx # type: ignore[misc] + self._m_idx = cutlass.Int32(m_idx) + self._n_idx = cutlass.Int32(n_idx) + self._l_idx = cutlass.Int32(l_idx) + self._is_valid_tile = cutlass.Boolean(is_valid_tile) + + def __extract_mlir_values__(self) -> list: + return [ + self._m_idx.ir_value(), + self._n_idx.ir_value(), + self._l_idx.ir_value(), + self._is_valid_tile.ir_value(), + ] + def __new_from_mlir_values__(self, values: list) -> "WorkTileInfo": + assert len(values) == 4 + return WorkTileInfo( + ( + cutlass.Int32(values[0]), + cutlass.Int32(values[1]), + cutlass.Int32(values[2]), + ), + cutlass.Boolean(values[3]), + ) -@cute.jit -def _work_tile_info_tile_idx_scalar(self: WorkTileInfo) -> cute.Coord: - return (self._m_idx, self._n_idx, self._l_idx) # type: ignore[attr-defined] + @property + @cute.jit + def tile_idx(self) -> cute.Coord: + return (self._m_idx, self._n_idx, self._l_idx) + @property + @cute.jit + def is_valid_tile(self) -> cutlass.Boolean: + return self._is_valid_tile -@cute.jit -def _work_tile_info_is_valid_tile_scalar(self: WorkTileInfo) -> cutlass.Boolean: - return self._is_valid_tile + @cute.jit + def update_from(self, other: _BaseWorkTileInfo) -> None: + m_idx, n_idx, l_idx = other.tile_idx + self._m_idx = cutlass.Int32(m_idx) + self._n_idx = cutlass.Int32(n_idx) + self._l_idx = cutlass.Int32(l_idx) + self._is_valid_tile = cutlass.Boolean(other.is_valid_tile) @cute.jit -def _work_tile_info_update_from_scalar(self: WorkTileInfo, other: WorkTileInfo) -> None: - m_idx, n_idx, l_idx = other.tile_idx - self._m_idx = cutlass.Int32(m_idx) # type: ignore[attr-defined] - self._n_idx = cutlass.Int32(n_idx) # type: ignore[attr-defined] - self._l_idx = cutlass.Int32(l_idx) # type: ignore[attr-defined] - self._is_valid_tile = cutlass.Boolean(other.is_valid_tile) - - -WorkTileInfo.__init__ = _work_tile_info_init_scalar # type: ignore[method-assign] -WorkTileInfo.__extract_mlir_values__ = _work_tile_info_extract_scalar # type: ignore[method-assign] -WorkTileInfo.__new_from_mlir_values__ = _work_tile_info_new_from_mlir_values_scalar # type: ignore[method-assign] -WorkTileInfo.tile_idx = property(_work_tile_info_tile_idx_scalar) # type: ignore[method-assign] -WorkTileInfo.is_valid_tile = property(_work_tile_info_is_valid_tile_scalar) # type: ignore[method-assign] -WorkTileInfo.update_from = _work_tile_info_update_from_scalar # type: ignore[attr-defined] - -# The scalarized methods are installed on CUTLASS' WorkTileInfo class. Some -# IR reconstruction paths resolve patched class methods in the owner module's -# globals rather than this module's globals; make the CuTe name used by the -# patches available there too. -_static_persistent_tile_scheduler.cute = cute +def _to_ts_work_tile_info(work_tile: _BaseWorkTileInfo) -> WorkTileInfo: + return WorkTileInfo(work_tile.tile_idx, work_tile.is_valid_tile) + # Variable-flow model # ------------------- @@ -730,10 +719,20 @@ def create_tma_async_pipeline_cfg( advance_on_wait : bool, optional Setting advance_on_wait=True on a PipelineConfig decouples the consumer's wait cursor from its release cursor by introducing two independent pipeline state objects: - | State | Advances when | Points to | - | ------------------------ | ----------------- | ------------------------------- | - | `consumer_state` | `ConsumerWait` | the *next* buffer to wait on | - | `consumer_release_state` | `ConsumerRelease` | the *oldest* un-released buffer | + + .. list-table:: + :header-rows: 1 + + * - State + - Advances when + - Points to + * - ``consumer_state`` + - ``ConsumerWait`` + - the *next* buffer to wait on + * - ``consumer_release_state`` + - ``ConsumerRelease`` + - the *oldest* un-released buffer + With this split, the consumer can wait on buffer N+1 (advancing consumer_state) while still holding buffer N (not yet released via consumer_release_state). interleave_stride : int or tuple[int, int, int, int], optional @@ -3469,7 +3468,7 @@ def create(self) -> None: @cute.jit def _make_initial_work_tile(self) -> WorkTileInfo: assert self.tile_scheduler is not None - return self.tile_scheduler.initial_work_tile_info() + return _to_ts_work_tile_info(self.tile_scheduler.initial_work_tile_info()) @cute.jit def initial_work_tile_info(self) -> WorkTileInfo: @@ -3533,7 +3532,7 @@ def _get_and_advance_work_tile_impl( ): assert self.tile_scheduler is not None self.tile_scheduler.advance_to_next_work() - work_tile = self.tile_scheduler.get_current_work() + work_tile = _to_ts_work_tile_info(self.tile_scheduler.get_current_work()) else: # CLC dynamic: read from the per-stage response buffer. stage_response_ptr = self._get_stage_response_ptr(stage_info.stage_idx) diff --git a/python/CuTeDSL/cutlass/experimental/task_scheduling/task.py b/python/CuTeDSL/cutlass/experimental/task_scheduling/task.py index cf1ff5f5d5..3800379688 100644 --- a/python/CuTeDSL/cutlass/experimental/task_scheduling/task.py +++ b/python/CuTeDSL/cutlass/experimental/task_scheduling/task.py @@ -43,9 +43,6 @@ from itertools import chain from typing import Any, Callable, FrozenSet, Optional, List, TYPE_CHECKING from typing import Tuple, cast -from cutlass.utils.static_persistent_tile_scheduler import ( - WorkTileInfo, -) if TYPE_CHECKING: from .schedule_builder import ScheduleResult @@ -68,6 +65,7 @@ StageInfo, MemoryResource, PipelineConfig, + WorkTileInfo, PipelineGroup, WorkQueue, SlotRouting, diff --git a/python/CuTeDSL/cutlass/jax/__init__.py b/python/CuTeDSL/cutlass/jax/__init__.py index 866cbd5b9e..a8a19776b5 100644 --- a/python/CuTeDSL/cutlass/jax/__init__.py +++ b/python/CuTeDSL/cutlass/jax/__init__.py @@ -10,6 +10,7 @@ # is strictly prohibited. from functools import cache +import importlib import logging import sys import warnings @@ -64,72 +65,112 @@ def is_available() -> bool: return True -if is_available(): - if sys.version_info[:2] < CUTE_DSL_JAX_MIN_SUPPORTED_PYTHON_VERSION: - minimum_version = ".".join(map(str, CUTE_DSL_JAX_MIN_SUPPORTED_PYTHON_VERSION)) - current_version = ".".join(map(str, sys.version_info[:3])) - warnings.warn( - f"cutlass.jax requires Python {minimum_version} or newer, but " - f"Python {current_version} is in use. Please upgrade Python.", - RuntimeWarning, - stacklevel=2, - ) - - from .primitive import cutlass_call - from .types import ( - jax_to_cutlass_dtype, - cutlass_to_jax_dtype, - jax_to_cutlass_layout_order, - cutlass_to_jax_layout_order, - from_dlpack, - JaxArray, - TensorSpec, - ) - from .compile import ( - release_compile_cache, - ) - from .ffi import ( - get_export_disabled_safety_checks, - find_cute_dsl_runtime_library, - register_ffi, - set_ffi_call_targets, - disable_automatic_ffi_registration, - is_ffi_registered, - get_cutlass_call_ffi_name, - get_cutlass_call_ffi_version, +@cache +def _warn_on_unsupported_python() -> None: + if sys.version_info[:2] >= CUTE_DSL_JAX_MIN_SUPPORTED_PYTHON_VERSION: + return + minimum_version = ".".join(map(str, CUTE_DSL_JAX_MIN_SUPPORTED_PYTHON_VERSION)) + current_version = ".".join(map(str, sys.version_info[:3])) + warnings.warn( + f"cutlass.jax requires Python {minimum_version} or newer, but " + f"Python {current_version} is in use. Please upgrade Python.", + RuntimeWarning, + stacklevel=2, ) - from . import testing + +# Export name -> (submodule, attribute). A None attribute yields the submodule +# itself. Every one of these submodules imports JAX at module scope, so binding +# them here rather than above is what keeps JAX out of the import. +_LAZY_ATTRS = { + "cutlass_call": (".primitive", "cutlass_call"), + "jax_to_cutlass_dtype": (".types", "jax_to_cutlass_dtype"), + "cutlass_to_jax_dtype": (".types", "cutlass_to_jax_dtype"), + "jax_to_cutlass_layout_order": (".types", "jax_to_cutlass_layout_order"), + "cutlass_to_jax_layout_order": (".types", "cutlass_to_jax_layout_order"), + "from_dlpack": (".types", "from_dlpack"), + "JaxArray": (".types", "JaxArray"), + "TensorSpec": (".types", "TensorSpec"), # This is a legacy name for TensorSpec. It will be removed eventually. - TensorMode = TensorSpec - - __all__ = [ - "CUTE_DSL_MIN_SUPPORTED_JAX_VERSION", - "CUTE_DSL_JAX_MIN_SUPPORTED_PYTHON_VERSION", - "cutlass_call", - "jax_to_cutlass_dtype", - "cutlass_to_jax_dtype", - "jax_to_cutlass_layout_order", - "cutlass_to_jax_layout_order", - "from_dlpack", - "JaxArray", - "TensorSpec", - "TensorMode", - "release_compile_cache", - "get_export_disabled_safety_checks", - "is_ffi_registered", - "register_ffi", - "set_ffi_call_targets", + "TensorMode": (".types", "TensorSpec"), + "release_compile_cache": (".compile", "release_compile_cache"), + "get_export_disabled_safety_checks": (".ffi", "get_export_disabled_safety_checks"), + "find_cute_dsl_runtime_library": (".ffi", "find_cute_dsl_runtime_library"), + "register_ffi": (".ffi", "register_ffi"), + "set_ffi_call_targets": (".ffi", "set_ffi_call_targets"), + "disable_automatic_ffi_registration": ( + ".ffi", "disable_automatic_ffi_registration", - "get_cutlass_call_ffi_name", - "get_cutlass_call_ffi_version", - "is_available", - "testing", - ] -else: - # export is_available check for callers or tests. - __all__ = [ - "CUTE_DSL_MIN_SUPPORTED_JAX_VERSION", - "CUTE_DSL_JAX_MIN_SUPPORTED_PYTHON_VERSION", - "is_available", - ] + ), + "is_ffi_registered": (".ffi", "is_ffi_registered"), + "get_cutlass_call_ffi_name": (".ffi", "get_cutlass_call_ffi_name"), + "get_cutlass_call_ffi_version": (".ffi", "get_cutlass_call_ffi_version"), + "testing": (".testing", None), +} + +_ALL_WITHOUT_JAX = [ + "CUTE_DSL_MIN_SUPPORTED_JAX_VERSION", + "CUTE_DSL_JAX_MIN_SUPPORTED_PYTHON_VERSION", + "is_available", +] + +_ALL_WITH_JAX = [ + "CUTE_DSL_MIN_SUPPORTED_JAX_VERSION", + "CUTE_DSL_JAX_MIN_SUPPORTED_PYTHON_VERSION", + "cutlass_call", + "jax_to_cutlass_dtype", + "cutlass_to_jax_dtype", + "jax_to_cutlass_layout_order", + "cutlass_to_jax_layout_order", + "from_dlpack", + "JaxArray", + "TensorSpec", + "TensorMode", + "release_compile_cache", + "get_export_disabled_safety_checks", + "is_ffi_registered", + "register_ffi", + "set_ffi_call_targets", + "disable_automatic_ffi_registration", + "get_cutlass_call_ffi_name", + "get_cutlass_call_ffi_version", + "is_available", + "testing", +] + + +def __getattr__(name: str): + """Bind the JAX-backed exports on first use. + + Importing any of them loads JAX and jaxlib. jaxlib's native extension + imports NumPy from C during module exec, and jaxlib builds that abort + rather than raise on a failed import kill the process before + ``is_available()`` can report False. Deferring the binding keeps both + ``import cutlass`` and ``import cutlass.jax`` independent of JAX, and costs + nothing for callers that never use the integration. + + ``__all__`` is served here too so that it keeps reflecting whether the + extensions are actually usable, as it did when it was assigned at import. + """ + if name == "__all__": + return list(_ALL_WITH_JAX if is_available() else _ALL_WITHOUT_JAX) + + target = _LAZY_ATTRS.get(name) + if target is None: + raise AttributeError(f"module {__name__!r} has no attribute {name!r}") + if not is_available(): + raise AttributeError( + f"module {__name__!r} has no attribute {name!r}: the CuTeDSL JAX " + "extensions are not available" + ) + + _warn_on_unsupported_python() + module_name, attribute = target + module = importlib.import_module(module_name, __name__) + value = module if attribute is None else getattr(module, attribute) + globals()[name] = value + return value + + +def __dir__(): + return sorted(set(globals()) | set(_LAZY_ATTRS)) diff --git a/python/CuTeDSL/cutlass/memory/tmem.py b/python/CuTeDSL/cutlass/memory/tmem.py index 0fb087c076..61705cd35b 100644 --- a/python/CuTeDSL/cutlass/memory/tmem.py +++ b/python/CuTeDSL/cutlass/memory/tmem.py @@ -479,8 +479,16 @@ def allocate( assert self.check_valid_num_columns(num_columns), ( f"num_columns must be multiple of 32 and power of two, and between 0 and {self._max_tmem_columns}" ) - assert self._num_allocated_columns + num_columns <= self._max_tmem_columns, ( - f"total allocated columns must be less than or equal to {self._max_tmem_columns}" + # Inside warp-specialized control flow the bookkeeping counter is a + # runtime value, so check capacity with a runtime assert. The backend + # strips it unless CUTE_DSL_ENABLE_ASSERTIONS=True. + from cutlass.cute.testing import assert_ as runtime_assert + + runtime_assert( + self._num_allocated_columns + num_columns <= self._max_tmem_columns, + f"total allocated columns must be less than or equal to {self._max_tmem_columns}", + loc=loc, + ip=ip, ) warp_idx = cute.arch.warp_idx(loc=loc, ip=ip) @@ -607,8 +615,15 @@ def free( warp_idx = cute.arch.make_warp_uniform(warp_idx, loc=loc, ip=ip) _is_allocator_warp = warp_idx == self._allocator_warp_id - assert num_columns <= self._num_allocated_columns, ( - "num_columns must be less than or equal to num_allocated_columns" + # Same as allocate(): the counter is a runtime value inside + # warp-specialized control flow, so this is a runtime assert. + from cutlass.cute.testing import assert_ as runtime_assert + + runtime_assert( + num_columns <= self._num_allocated_columns, + "num_columns must be less than or equal to num_allocated_columns", + loc=loc, + ip=ip, ) if const_expr(num_columns != 0): assert self.check_valid_num_columns(num_columns), "num_columns is invalid" diff --git a/python/CuTeDSL/cutlass/pipeline/helpers.py b/python/CuTeDSL/cutlass/pipeline/helpers.py index 9c0f0cd956..3e0e8b7c8d 100644 --- a/python/CuTeDSL/cutlass/pipeline/helpers.py +++ b/python/CuTeDSL/cutlass/pipeline/helpers.py @@ -16,6 +16,7 @@ from typing import Any, Optional, Union, cast import warnings +import cutlass import cutlass.cute as cute from cutlass._mlir import ir from cutlass import base_dsl @@ -99,7 +100,21 @@ class CooperativeGroup: CooperativeGroup contains size restrictions for an Agent. """ - def __init__(self, agent: Agent, size: Union[int, Int32] = 1): + def __init__( + self, + agent: Agent, + size: Union[int, Int32] = 1, + alignment: Optional[int] = None, + ): + if alignment is not None: + warnings.warn( + "The 'alignment' parameter of CooperativeGroup's constructor is " + "deprecated and will be removed in a subsequent release, please " + "remove it from your code.", + DeprecationWarning, + stacklevel=2, + ) + if agent in [ Agent.Thread, Agent.Warp, @@ -867,10 +882,22 @@ def reset_count( @cute.jit def advance( self, + count: int | Int32 = 1, *, loc: Optional[ir.Location] = None, ip: Optional[ir.InsertionPoint] = None, ) -> None: + if cutlass.const_expr(isinstance(count, int) and count < 1): + raise ValueError("count must be >= 1") + + if cutlass.const_expr(isinstance(count, Int32) or count > 1): + self._count += count + self._index += count + phase_flips = self._index // self._stages + self._index = self._index % self._stages + self._phase ^= phase_flips % 2 + return + self._index += 1 self._count += 1 diff --git a/python/CuTeDSL/cutlass/pipeline/sm90.py b/python/CuTeDSL/cutlass/pipeline/sm90.py index 548b22c47c..b39f6a7912 100644 --- a/python/CuTeDSL/cutlass/pipeline/sm90.py +++ b/python/CuTeDSL/cutlass/pipeline/sm90.py @@ -1111,8 +1111,8 @@ def reset( self.__state.reset_count(loc=loc, ip=ip) def current_handle(self) -> ImmutableResourceHandle: - """Get the current handle for the producer.""" - return PipelineProducer.ImmutableResourceHandle(self.__pipeline, self.__state) + """Get the current immutable handle for the producer.""" + return PipelineProducer.ImmutableResourceHandle(self.__pipeline, self.__state.clone()) @dsl_user_op def acquire( @@ -1142,12 +1142,13 @@ def acquire( @dsl_user_op def advance( self, + count: int | Int32 = 1, *, loc: Optional[ir.Location] = None, ip: Optional[ir.InsertionPoint] = None, ) -> None: """Move to the next pipeline stage.""" - self.__state.advance(loc=loc, ip=ip) + self.__state.advance(count, loc=loc, ip=ip) @dsl_user_op def acquire_and_advance( @@ -1356,8 +1357,8 @@ def reset( self.__state.reset_count(loc=loc, ip=ip) def current_handle(self) -> ImmutableResourceHandle: - """Get the current handle for the consumer.""" - return PipelineConsumer.ImmutableResourceHandle(self.__pipeline, self.__state) + """Get the current immutable handle for the consumer.""" + return PipelineConsumer.ImmutableResourceHandle(self.__pipeline, self.__state.clone()) @dsl_user_op def wait( @@ -1386,6 +1387,7 @@ def wait( @dsl_user_op def advance( self, + count: int | Int32 = 1, *, loc: Optional[ir.Location] = None, ip: Optional[ir.InsertionPoint] = None, @@ -1395,7 +1397,7 @@ def advance( This updates the internal state to point to the next buffer in the pipeline. Should be called after consuming data from the current buffer. """ - self.__state.advance(loc=loc, ip=ip) + self.__state.advance(count, loc=loc, ip=ip) @dsl_user_op def wait_and_advance( diff --git a/python/CuTeDSL/cutlass/testing.py b/python/CuTeDSL/cutlass/testing.py index 5636a5e9ff..34a182a27a 100644 --- a/python/CuTeDSL/cutlass/testing.py +++ b/python/CuTeDSL/cutlass/testing.py @@ -290,6 +290,22 @@ def _does_kernel_use_stream( return False +def _does_kernel_use_stream_with_jit_retry( + kernel: Callable[..., Any], + stream: cuda_driver.CUstream, + args: tuple[Any, ...], + kwargs: dict[str, Any], +) -> bool: + uses_stream = _does_kernel_use_stream(kernel, stream, args, kwargs) + if uses_stream or not hasattr(kernel, "_dsl_cls"): + return uses_stream + + # The first invocation of an uncompiled @cute.jit function can spend the + # capture attempt compiling instead of recording a launch. Retry once with + # the now-compiled callable before reporting a stream mismatch. + return _does_kernel_use_stream(kernel, stream, args, kwargs) + + def benchmark( callable: Callable, *, @@ -407,7 +423,7 @@ def workspace_generator(): if ( not use_cuda_graphs and int(stream) != int(cuda_driver.CUstream_flags.CU_STREAM_DEFAULT) - and not _does_kernel_use_stream( + and not _does_kernel_use_stream_with_jit_retry( callable, stream, workspaces[0].args, workspaces[0].kwargs ) ): @@ -743,7 +759,9 @@ def _benchmark_for_autotune( if int(current_stream) != int( cuda_driver.CUstream(cuda_driver.CUstream_flags.CU_STREAM_DEFAULT) - ) and not _does_kernel_use_stream(callable, current_stream, args, kwargs): + ) and not _does_kernel_use_stream_with_jit_retry( + callable, current_stream, args, kwargs + ): raise ValueError(f"Incorrect stream passed to kernel: {current_stream}") if use_cold_l2: diff --git a/python/CuTeDSL/cutlass/utils/blackwell_helpers.py b/python/CuTeDSL/cutlass/utils/blackwell_helpers.py index f8c5f247a4..14cecfadf2 100644 --- a/python/CuTeDSL/cutlass/utils/blackwell_helpers.py +++ b/python/CuTeDSL/cutlass/utils/blackwell_helpers.py @@ -892,6 +892,61 @@ def make_smem_layout_b( ) +@dsl_user_op +def make_smem_layout_c( + tiled_mma: cute.TiledMma, + mma_tiler_mnk: cute.Tile, + c_dtype: Type[Numeric], + num_stages: int, + is_m_major: bool, + *, + loc: Optional[ir.Location] = None, + ip: Optional[ir.InsertionPoint] = None, +) -> Union[cute.Layout, cute.ComposedLayout]: + """This function helps with: + + 1. Get the partitioned shape of the C tensor based on the tiled_mma & MMA tiler. + 2. Select the heuristic SMEM layout atom based on the C tensor's majorness, the data type, and the major mode size. + 3. cute.Tile the SMEM layout atom to the MMA tile shape. + 4. Stage the SMEM layout based on the number of stages. + + Useful for when we don't care about epilogue subtiling, which get_smem_layout_epi is designed for. + + :param tiled_mma: The tiled MMA used to partition tensor C + :type tiled_mma: cute.TiledMma + :param mma_tiler_mnk: The MMA tile shape + :type mma_tiler_mnk: cute.cute.Tile + :param c_dtype: The element type for tensor C + :type c_dtype: Type[Numeric] + :param num_stages: The number of pipeline stages for tensor C + :type num_stages: int + + :return: SMEM layout for tensor C + :rtype: Union[cute.Layout, cute.ComposedLayout] + """ + + c_major_mode = OperandMajorMode.MN if is_m_major else OperandMajorMode.K + c_smem_shape = tiled_mma.partition_shape_C( + cute.dice(mma_tiler_mnk, (1, 1, None), loc=loc, ip=ip), loc=loc, ip=ip + ) + c_smem_shape_mn = ( + cute.size(c_smem_shape[0][0], loc=loc, ip=ip) * c_smem_shape[1], + cute.size(c_smem_shape[0][1], loc=loc, ip=ip) * c_smem_shape[2], + ) + smem_layout_atom_kind = get_smem_layout_atom_ab( + c_major_mode, c_dtype, c_smem_shape_mn, loc=loc, ip=ip + ) + c_smem_layout_atom = make_smem_layout_atom( + smem_layout_atom_kind, c_dtype, loc=loc, ip=ip + ) + + c_smem_shape = cute.append(c_smem_shape, num_stages, loc=loc, ip=ip) + order = (2, 1, 3) if is_m_major else (1, 2, 3) + return tile_to_mma_shape( + c_smem_layout_atom, c_smem_shape, order=order, loc=loc, ip=ip + ) + + @dsl_user_op def get_smem_layout_atom_epi( layout: LayoutEnum, diff --git a/python/CuTeDSL/cutlass/utils/mixed_input_helpers.py b/python/CuTeDSL/cutlass/utils/mixed_input_helpers.py index bdd94b0dcf..5b19981b43 100644 --- a/python/CuTeDSL/cutlass/utils/mixed_input_helpers.py +++ b/python/CuTeDSL/cutlass/utils/mixed_input_helpers.py @@ -1025,7 +1025,20 @@ def cvt_tensor_a( def cvt_tensor_a_mxf8(src: cute.Tensor, dtype: type[cutlass.Numeric]) -> cute.TensorSSA: - """Convert an int4 A-tile slice to mxfp8 via :func:`cute.arch.cvt_i4_mxf8_intrinsic`.""" + """Convert an int4 A-tile slice to mxfp8 via :func:`cute.arch.cvt_i4_mxf8_intrinsic `. + + The slice is loaded into a vector and converted in 8-element chunks, so its + element count must be a multiple of 8. The result keeps the source shape. + + :param src: A-tile slice of int4 values. + :type src: cute.Tensor + :param dtype: FP8 destination numeric type, typically ``cutlass.Float8E4M3FN``. + :type dtype: type[cutlass.Numeric] + :return: The converted tile with ``src``'s shape and ``dtype`` elements. + :rtype: cute.TensorSSA + :raises AssertionError: if the slice does not hold a multiple of 8 elements, + which the underlying intrinsic requires to process whole chunks. + """ rst = src.load() return cute.TensorSSA.from_vector( cute.arch.cvt_i4_mxf8_intrinsic(rst, cute.size(rst.shape), dtype), diff --git a/python/CuTeDSL/requirements-cu13.txt b/python/CuTeDSL/requirements-cu13.txt index ec4a195b35..aa8aa96826 100644 --- a/python/CuTeDSL/requirements-cu13.txt +++ b/python/CuTeDSL/requirements-cu13.txt @@ -1,3 +1,3 @@ # Use `pip install -r requirements-cu13.txt` with the present file to install a # wheel consistent with the present state of the github repository -nvidia-cutlass-dsl[cu13]==4.8.0.dev0 +nvidia-cutlass-dsl[cu13]==4.8.0 diff --git a/python/CuTeDSL/requirements.txt b/python/CuTeDSL/requirements.txt index 0bcf40a597..f03bb49c22 100644 --- a/python/CuTeDSL/requirements.txt +++ b/python/CuTeDSL/requirements.txt @@ -1,3 +1,3 @@ # Use `pip install -r requirements.txt` with the present file to install a # wheel consistent with the present state of the github repository -nvidia-cutlass-dsl==4.8.0.dev0 +nvidia-cutlass-dsl==4.8.0 diff --git a/python/cutlass_library/emit_kernel_listing.py b/python/cutlass_library/emit_kernel_listing.py index 3d593f58ec..7ec5bdd5e3 100755 --- a/python/cutlass_library/emit_kernel_listing.py +++ b/python/cutlass_library/emit_kernel_listing.py @@ -336,6 +336,18 @@ def emit_gemm_kernel_testlist(manifest, curr_build_dir, arch, mode sm100_mma_filter_regex_1sm_runtime = "cutlass3x_sm100_tensorop.*(" + ").*(".join([ "|".join(x) for x in [sm100_mma_data_type_runtime_dtype, sm100_mma_cluster_size, sm100_mma_layouts]]) + ").*1sm.*" sm100_mma_filter_regex_2sm_runtime = "cutlass3x_sm100_tensorop.*(" + ").*(".join([ "|".join(x) for x in [sm100_mma_data_type_runtime_dtype, sm100_mma_cluster_size, sm100_mma_layouts]]) + ").*2sm.*" + sm107_mma_cluster_size = [ + '2x2x1', + '0x0x1' # dynamic cluster + ] + + sm107_mma_data_type_general = [ + "gemm_f8_f8_f32_f16_f16", + ] + + sm107_mma_filter_regex_1sm = "cutlass3x_sm107_tensorop.*(" + ").*(".join([ "|".join(x) for x in [sm107_mma_data_type_general, sm107_mma_cluster_size, sm100_mma_layouts]]) + ").*_breuse_1sm.*" + sm107_mma_filter_regex_2sm = "cutlass3x_sm107_tensorop.*(" + ").*(".join([ "|".join(x) for x in [sm107_mma_data_type_general, sm107_mma_cluster_size, sm100_mma_layouts]]) + ").*_breuse_2sm.*" + # # Block Scale Gemm # @@ -367,11 +379,28 @@ def emit_gemm_kernel_testlist(manifest, curr_build_dir, arch, mode # regex list must be in kernel procedural name order block_scaled_filter_regex_1sm = "cutlass3x_sm100_bstensorop.*(" + ").*(".join([ "|".join(x) for x in [block_scaled_data_type, block_scaled_tile_k, block_scaled_cluster_size, block_scaled_layouts]]) + ").*1sm.*" block_scaled_filter_regex_2sm = "cutlass3x_sm100_bstensorop.*(" + ").*(".join([ "|".join(x) for x in [block_scaled_data_type, block_scaled_tile_k, block_scaled_cluster_size, block_scaled_layouts]]) + ").*2sm.*" - + sm103_block_scaled_prefetch_policy = ['tmapf'] sm103_block_scaled_filter_regex_1sm = "cutlass3x_sm103_bstensorop.*(" + ").*(".join([ "|".join(x) for x in [sm103_block_scaled_data_type, sm103_block_scaled_tile_k, block_scaled_cluster_size, block_scaled_layouts]]) + ").*1sm.*(" + "|".join(sm103_block_scaled_prefetch_policy) + ").*" sm103_block_scaled_filter_regex_2sm = "cutlass3x_sm103_bstensorop.*(" + ").*(".join([ "|".join(x) for x in [sm103_block_scaled_data_type, sm103_block_scaled_tile_k, block_scaled_cluster_size, block_scaled_layouts]]) + ").*2sm.*(" + "|".join(sm103_block_scaled_prefetch_policy) + ").*" + sm107_mma_bs_data_type_general = [ + "gemm_ue8m0xf8_ue8m0xf8_f32_f16_f16", + "gemm_ue8m0xf8_ue8m0xf8_f32_bf16_e5m2", + ] + + sm107_mma_bs_filter_regex_1sm = "cutlass3x_sm107_bstensorop.*(" + ").*(".join([ "|".join(x) for x in [sm107_mma_bs_data_type_general, sm107_mma_cluster_size, block_scaled_layouts]]) + ").*_breuse_1sm.*" + sm107_mma_bs_filter_regex_2sm = "cutlass3x_sm107_bstensorop.*(" + ").*(".join([ "|".join(x) for x in [sm107_mma_bs_data_type_general, sm107_mma_cluster_size, block_scaled_layouts]]) + ").*_breuse_2sm.*" + + sm107_mma_nvf4_data_type_general = [ + "gemm_ue8m0xe2m1_ue8m0xe2m1_f32_f16_f16", + "gemm_ue4m3xe2m1_ue4m3xe2m1_f32_bf16_bf16", + "gemm_ue5m3xe2m1_ue5m3xe2m1_f32_f16_f16", + ] + + sm107_mma_nvf4_filter_regex_1sm = "cutlass3x_sm107_bstensorop.*(" + ").*(".join([ "|".join(x) for x in [sm107_mma_nvf4_data_type_general, sm107_mma_cluster_size, block_scaled_layouts]]) + ").*_breuse_1sm.*" + sm107_mma_nvf4_filter_regex_2sm = "cutlass3x_sm107_bstensorop.*(" + ").*(".join([ "|".join(x) for x in [sm107_mma_nvf4_data_type_general, sm107_mma_cluster_size, block_scaled_layouts]]) + ").*_breuse_2sm.*" + if arch in ["100a", "100f"]: kernel_filter = f"({sm100_mma_filter_regex_1sm})|" \ f"({sm100_mma_filter_regex_2sm})|" \ diff --git a/python/cutlass_library/gemm_operation.py b/python/cutlass_library/gemm_operation.py index 0229bcaed7..77ab2953c6 100644 --- a/python/cutlass_library/gemm_operation.py +++ b/python/cutlass_library/gemm_operation.py @@ -355,6 +355,22 @@ def get_collective_tile_shape(self): if opcode_class_main in [OpcodeClass.TensorOp, OpcodeClass.BlockScaledTensorOp, OpcodeClass.SparseTensorOp, OpcodeClass.BlockScaledSparseTensorOp]: tile_shape_m = instruction_shape[0] tile_shape_n = instruction_shape[1] + + _sm107_breuse_schedules = { + KernelScheduleType.TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithBreuse, + KernelScheduleType.TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithBreuse, + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse, + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse, + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse, + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse, + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse, + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse, + } + # Re-adjusting the tile shape m-mode, since it can be double the instruction's + # m-mode if b-reuse is enabled + if self.kernel_schedule in _sm107_breuse_schedules: + tile_shape_m = tile_shape_m * 2 + return (tile_shape_m, tile_shape_n, tile_shape_k) # Generates the full kernel function name diff --git a/python/cutlass_library/generator.py b/python/cutlass_library/generator.py index e2d1309834..d873ed58f0 100644 --- a/python/cutlass_library/generator.py +++ b/python/cutlass_library/generator.py @@ -6781,6 +6781,21 @@ def GenerateSM90_TensorOp_1684_symm_complex_gaussian(manifest, cuda_version): get_pruning_level_from_global_level ) +# Rubin SM 107 generators + +try: + from cutlass_library.sm107_utils import ( + generate_f8f6f4_math_instructions_sm107, + generate_mxf8f6f4_math_instructions_sm107, + generate_mxnvf4_math_instructions_sm107, + ) +except ImportError: + from sm107_utils import ( + generate_f8f6f4_math_instructions_sm107, + generate_mxf8f6f4_math_instructions_sm107, + generate_mxnvf4_math_instructions_sm107, + ) + ################################################################################################### def get_tma_alignment_elt(data_type : DataType, is_f8f6f4 : bool = True ) -> int: @@ -12057,6 +12072,403 @@ def tile_schedulers(kernel_schedule): tile_schedulers = tile_schedulers(kernel_schedule), gemm_kind = gemm_kind) + +def GenerateSM107_TensorOp_fp8_UMMA_gemm(manifest, cuda_version): + if not CudaToolkitVersionSatisfies(cuda_version, 13, 3): + return + + instantiation_level = manifest.get_instantiation_level(pruned_level=591, default_level=591, exhaustive_level=9999) + + layouts = [ + [[LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 0]], + [[LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 0]], + [[LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 0]], + [[LayoutType.RowMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 0]], + [[LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 0]], + [[LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.RowMajor, 0]], + [[LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 0]], + [[LayoutType.RowMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.RowMajor, 0]], + ] + + min_cc = 107 + max_cc = 107 + + epi_type = DataType.f32 + + math_instructions_1sm, math_instructions_2sm = generate_f8f6f4_math_instructions_sm107(instantiation_level) + + cluster_shapes_1sm, cluster_shapes_2sm = generate_cluster_shapes_sm100(instantiation_level) + + tile_schedulers = [TileSchedulerType.Default] + + cd_data_types = [ + {"c_type": DataType.f16, "d_type": DataType.f16 }, + {"c_type": DataType.f16, "d_type": DataType.e4m3}, + {"c_type": DataType.f16, "d_type": DataType.e5m2}, + {"c_type": DataType.bf16, "d_type": DataType.bf16}, + {"c_type": DataType.bf16, "d_type": DataType.e4m3}, + {"c_type": DataType.bf16, "d_type": DataType.e5m2}, + {"c_type": DataType.f32, "d_type": DataType.f32 }, + # Void-C kernels + {"c_type": DataType.void, "d_type": DataType.f16 }, + {"c_type": DataType.void, "d_type": DataType.bf16}, + {"c_type": DataType.void, "d_type": DataType.f32 }, + {"c_type": DataType.void, "d_type": DataType.e4m3}, + {"c_type": DataType.void, "d_type": DataType.e5m2}, + ] + + for b_reuse in [False, True]: + kernel_schedule_1sm = (KernelScheduleType.TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithBreuse + if b_reuse else + KernelScheduleType.TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithoutBreuse) + kernel_schedule_2sm = (KernelScheduleType.TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithBreuse + if b_reuse else + KernelScheduleType.TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithoutBreuse) + + # 1SM MMA kernels + for math_inst in math_instructions_1sm: + tile_descriptions = [] + for cluster_shape in cluster_shapes_1sm: + multiplier_1sm = (1, 1, 1) if cluster_shape == DynamicClusterShape else cluster_shape + m_multiplier = 2 if b_reuse else 1 + tile_descriptions.append( + TileDescription([ + math_inst.instruction_shape[0] * m_multiplier * multiplier_1sm[0], + math_inst.instruction_shape[1] * multiplier_1sm[1], + math_inst.instruction_shape[2] * 2 * multiplier_1sm[2]], + 0, [2, 1, 1], math_inst, min_cc, max_cc, cluster_shape)) + + for layout in layouts: + layout[2][1] = 128 // DataTypeSize[DataType.f16] # use f16 as reference alignment + for tile_description in tile_descriptions: + for layout in layouts: + for cd in cd_data_types: + layout[2][1] = 128 // DataTypeSize[cd["d_type"]] + if layout[1][0] == LayoutType.RowMajor and tile_description.math_instruction.instruction_shape[1] % 16 != 0: + continue + data_type = { + "a_type": math_inst.element_a, + "b_type": math_inst.element_b, + "c_type": cd["c_type"], + "d_type": cd["d_type"], + "acc_type": math_inst.element_accumulator, + "epi_type": epi_type, + } + ops = CreateGemmUniversal3xOperator(manifest, [layout], [tile_description], data_type, + [[kernel_schedule_1sm, EpilogueScheduleType.TmaWarpSpecialized1Sm]], + tile_schedulers=tile_schedulers) + + # 2SM MMA kernels + for math_inst in math_instructions_2sm: + tile_descriptions = [] + for cluster_shape in cluster_shapes_2sm: + multiplier_2sm = (1, 1, 1) if cluster_shape == DynamicClusterShape else (cluster_shape[0] // 2, cluster_shape[1], cluster_shape[2]) + m_multiplier = 2 if b_reuse else 1 + tile_descriptions.append( + TileDescription([ + math_inst.instruction_shape[0] * m_multiplier * multiplier_2sm[0], + math_inst.instruction_shape[1] * multiplier_2sm[1], + math_inst.instruction_shape[2] * 2 * multiplier_2sm[2]], + 0, [2, 1, 1], math_inst, min_cc, max_cc, cluster_shape)) + + for cd in cd_data_types: + for layout in layouts: + layout[2][1] = 128 // DataTypeSize[cd["d_type"]] + + # Similar to SM100, for SM107 2SM TMA epilogue with 16-bit output requires CtaN divisible + # by 64 when CtaN > 128. N=160 and N=224 fail the upcast<64> stride-divisibility check + # in the SMEM swizzle layout. Skip tile shapes where CtaN is not 64-aligned. + if DataTypeSize[cd["d_type"]] == 16: + tile_descs_for_dtype = [ + td for td in tile_descriptions if td.threadblock_shape[1] <= 128 or td.threadblock_shape[1] % 64 == 0 + ] + else: + tile_descs_for_dtype = tile_descriptions + if not tile_descs_for_dtype: + continue + + data_type = { + "a_type": math_inst.element_a, + "b_type": math_inst.element_b, + "c_type": cd["c_type"], + "d_type": cd["d_type"], + "acc_type": math_inst.element_accumulator, + "epi_type": epi_type, + } + ops = CreateGemmUniversal3xOperator(manifest, layouts, tile_descs_for_dtype, data_type, + [[kernel_schedule_2sm, EpilogueScheduleType.TmaWarpSpecialized2Sm]], + tile_schedulers=tile_schedulers) + + +def GenerateSM107_TensorOp_mxf8f6f4_UMMA_gemm_with_block_scaled( + manifest, cuda_version, gemm_kind=GemmKind.BlockScaledUniversal3x, +): + if not CudaToolkitVersionSatisfies(cuda_version, 13, 3): + return + + instantiation_level = manifest.get_instantiation_level(pruned_level=591, default_level=591, exhaustive_level=9999) + + layouts = [ + [[LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 0]], + [[LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 0]], + [[LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 0]], + [[LayoutType.RowMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 0]], + [[LayoutType.ColumnMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 0]], + [[LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.RowMajor, 0]], + [[LayoutType.RowMajor, 16], [LayoutType.ColumnMajor, 16], [LayoutType.RowMajor, 0]], + [[LayoutType.RowMajor, 16], [LayoutType.RowMajor, 16], [LayoutType.RowMajor, 0]], + ] + min_cc = 107 + max_cc = 107 + + epi_type = DataType.f32 + + # Only a/b = f8 (runtime dtype) is supported for SM107 blockscaled today. + math_instructions_1sm, math_instructions_2sm = generate_mxf8f6f4_math_instructions_sm107(instantiation_level) + + cluster_shapes_1sm, cluster_shapes_2sm = generate_cluster_shapes_sm100(instantiation_level) + + # No stream-K for SM107 blockscaled -- default tile scheduler only. + tile_schedulers = [TileSchedulerType.Default] + + cd_data_types = [ + {"c_type": DataType.f16, "d_type": DataType.f16 }, + {"c_type": DataType.f16, "d_type": DataType.e4m3}, + {"c_type": DataType.f16, "d_type": DataType.e5m2}, + {"c_type": DataType.bf16, "d_type": DataType.bf16}, + {"c_type": DataType.bf16, "d_type": DataType.e4m3}, + {"c_type": DataType.bf16, "d_type": DataType.e5m2}, + {"c_type": DataType.f32, "d_type": DataType.f32 }, + {"c_type": DataType.void, "d_type": DataType.f16 }, + {"c_type": DataType.void, "d_type": DataType.bf16}, + {"c_type": DataType.void, "d_type": DataType.f32 }, + {"c_type": DataType.void, "d_type": DataType.e4m3}, + {"c_type": DataType.void, "d_type": DataType.e5m2}, + ] + epilogue_schedule_1sm = EpilogueScheduleType.TmaWarpSpecialized1Sm + epilogue_schedule_2sm = EpilogueScheduleType.TmaWarpSpecialized2Sm + def make_data_type(math_inst, cd): + return { + "a_type": math_inst.element_a, + "b_type": math_inst.element_b, + "c_type": cd["c_type"], + "d_type": cd["d_type"], + "acc_type": math_inst.element_accumulator, + "epi_type": epi_type, + "sf_type": math_inst.element_scale_factor, + "sfd_type": {"type": DataType.void, "vector_size": None, "layout": None}, + } + + for b_reuse in [False, True]: + kernel_schedule_1sm = (KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse + if b_reuse else + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse) + kernel_schedule_2sm = (KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse + if b_reuse else + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse) + # 1SM MMA kernels + for math_inst in math_instructions_1sm: + assert math_inst.opcode_class == OpcodeClass.BlockScaledTensorOp + tile_descriptions = [] + for cluster_shape in cluster_shapes_1sm: + multiplier_1sm = (1, 1, 1) if cluster_shape == DynamicClusterShape else cluster_shape + m_multiplier = 2 if b_reuse else 1 + tile_descriptions.append( + TileDescription([ + math_inst.instruction_shape[0] * m_multiplier * multiplier_1sm[0], + math_inst.instruction_shape[1] * multiplier_1sm[1], + math_inst.instruction_shape[2] * 2 * multiplier_1sm[2]], + 0, [2, 1, 1], math_inst, min_cc, max_cc, cluster_shape)) + + for layout in layouts: + layout[2][1] = 128 // DataTypeSize[DataType.f16] # use f16 as reference alignment + for tile_description in tile_descriptions: + for layout in layouts: + for cd in cd_data_types: + layout[2][1] = 128 // DataTypeSize[cd["d_type"]] + if layout[1][0] == LayoutType.RowMajor and tile_description.math_instruction.instruction_shape[1] % 16 != 0: + continue + data_type = make_data_type(math_inst, cd) + ops = CreateGemmUniversal3xOperator(manifest, [layout], [tile_description], data_type, + [[kernel_schedule_1sm, epilogue_schedule_1sm]], + tile_schedulers=tile_schedulers, gemm_kind=gemm_kind) + + # 2SM MMA kernels + for math_inst in math_instructions_2sm: + assert math_inst.opcode_class == OpcodeClass.BlockScaledTensorOp + tile_descriptions = [] + for cluster_shape in cluster_shapes_2sm: + multiplier_2sm = (1, 1, 1) if cluster_shape == DynamicClusterShape else (cluster_shape[0] // 2, cluster_shape[1], cluster_shape[2]) + m_multiplier = 2 if b_reuse else 1 + tile_descriptions.append( + TileDescription([ + math_inst.instruction_shape[0] * m_multiplier * multiplier_2sm[0], + math_inst.instruction_shape[1] * multiplier_2sm[1], + math_inst.instruction_shape[2] * 2 * multiplier_2sm[2]], + 0, [2, 1, 1], math_inst, min_cc, max_cc, cluster_shape)) + + for cd in cd_data_types: + for layout in layouts: + layout[2][1] = 128 // DataTypeSize[cd["d_type"]] + + # Similar to SM100, for SM107 2SM TMA epilogue with 16-bit output requires CtaN divisible + # by 64 when CtaN > 128. N=160 and N=224 fail the upcast<64> stride-divisibility check + # in the SMEM swizzle layout. Skip tile shapes where CtaN is not 64-aligned. + if DataTypeSize[cd["d_type"]] == 16: + tile_descs_for_dtype = [ + td for td in tile_descriptions if td.threadblock_shape[1] <= 128 or td.threadblock_shape[1] % 64 == 0 + ] + else: + tile_descs_for_dtype = tile_descriptions + if not tile_descs_for_dtype: + continue + + data_type = make_data_type(math_inst, cd) + ops = CreateGemmUniversal3xOperator(manifest, layouts, tile_descs_for_dtype, data_type, + [[kernel_schedule_2sm, epilogue_schedule_2sm]], + tile_schedulers=tile_schedulers, gemm_kind=gemm_kind) + + +def GenerateSM107_TensorOp_mxnvf4_UMMA_gemm_with_block_scaled( + manifest, cuda_version, gemm_kind=GemmKind.BlockScaledUniversal3x, +): + if not CudaToolkitVersionSatisfies(cuda_version, 13, 3): + return + + instantiation_level = manifest.get_instantiation_level(pruned_level=591, default_level=591, exhaustive_level=9999) + + # A/B = e2m1 (compile-time dtype -- NVF4 has no other 4-bit element), so reference + # alignment is 128 // DataTypeSize[e2m1] = 32. + # MXNV_F4 only supports RowMajor A and ColumnMajor B + layouts = [ + [[LayoutType.RowMajor, 32], [LayoutType.ColumnMajor, 32], [LayoutType.ColumnMajor, 0]], + [[LayoutType.RowMajor, 32], [LayoutType.ColumnMajor, 32], [LayoutType.RowMajor, 0]], + ] + min_cc = 107 + max_cc = 107 + + epi_type = DataType.f32 + + # a/b are the compile-time e2m1 type for SM107 blockscaled NVF4. The generated math + # instructions don't differ by scale-factor vector size, so the same lists are reused + # below for both the Vs16 and Vs32 kernel schedules. + math_instructions_1sm, math_instructions_2sm = generate_mxnvf4_math_instructions_sm107(instantiation_level) + + cluster_shapes_1sm, cluster_shapes_2sm = generate_cluster_shapes_sm100(instantiation_level) + + # No stream-K for SM107 blockscaled -- default tile scheduler only. + tile_schedulers = [TileSchedulerType.Default] + + cd_data_types = [ + {"c_type": DataType.f16, "d_type": DataType.f16 }, + {"c_type": DataType.f16, "d_type": DataType.e4m3}, + {"c_type": DataType.f16, "d_type": DataType.e5m2}, + {"c_type": DataType.bf16, "d_type": DataType.bf16}, + {"c_type": DataType.bf16, "d_type": DataType.e4m3}, + {"c_type": DataType.bf16, "d_type": DataType.e5m2}, + {"c_type": DataType.f32, "d_type": DataType.f32 }, + {"c_type": DataType.void, "d_type": DataType.f16 }, + {"c_type": DataType.void, "d_type": DataType.bf16}, + {"c_type": DataType.void, "d_type": DataType.f32 }, + {"c_type": DataType.void, "d_type": DataType.e4m3}, + {"c_type": DataType.void, "d_type": DataType.e5m2}, + ] + epilogue_schedule_1sm = EpilogueScheduleType.TmaWarpSpecialized1Sm + epilogue_schedule_2sm = EpilogueScheduleType.TmaWarpSpecialized2Sm + def make_data_type(math_inst, cd): + return { + "a_type": math_inst.element_a, + "b_type": math_inst.element_b, + "c_type": cd["c_type"], + "d_type": cd["d_type"], + "acc_type": math_inst.element_accumulator, + "epi_type": epi_type, + "sf_type": math_inst.element_scale_factor, + "sfd_type": {"type": DataType.void, "vector_size": None, "layout": None}, + } + + for vector_size in [16, 32]: + for b_reuse in [False, True]: + if vector_size == 16: + kernel_schedule_1sm = (KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse + if b_reuse else + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse) + kernel_schedule_2sm = (KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse + if b_reuse else + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse) + else: + kernel_schedule_1sm = (KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse + if b_reuse else + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse) + kernel_schedule_2sm = (KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse + if b_reuse else + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse) + # 1SM MMA kernels + for math_inst in math_instructions_1sm: + assert math_inst.opcode_class == OpcodeClass.BlockScaledTensorOp + tile_descriptions = [] + for cluster_shape in cluster_shapes_1sm: + multiplier_1sm = (1, 1, 1) if cluster_shape == DynamicClusterShape else cluster_shape + m_multiplier = 2 if b_reuse else 1 + tile_descriptions.append( + TileDescription([ + math_inst.instruction_shape[0] * m_multiplier * multiplier_1sm[0], + math_inst.instruction_shape[1] * multiplier_1sm[1], + math_inst.instruction_shape[2] * 2 * multiplier_1sm[2]], + 0, [2, 1, 1], math_inst, min_cc, max_cc, cluster_shape)) + + for layout in layouts: + layout[2][1] = 128 // DataTypeSize[DataType.f16] # use f16 as reference alignment + for tile_description in tile_descriptions: + for layout in layouts: + for cd in cd_data_types: + layout[2][1] = 128 // DataTypeSize[cd["d_type"]] + if layout[1][0] == LayoutType.RowMajor and tile_description.math_instruction.instruction_shape[1] % 16 != 0: + continue + data_type = make_data_type(math_inst, cd) + ops = CreateGemmUniversal3xOperator(manifest, [layout], [tile_description], data_type, + [[kernel_schedule_1sm, epilogue_schedule_1sm]], + tile_schedulers=tile_schedulers, gemm_kind=gemm_kind) + + # 2SM MMA kernels + for math_inst in math_instructions_2sm: + assert math_inst.opcode_class == OpcodeClass.BlockScaledTensorOp + tile_descriptions = [] + for cluster_shape in cluster_shapes_2sm: + multiplier_2sm = (1, 1, 1) if cluster_shape == DynamicClusterShape else (cluster_shape[0] // 2, cluster_shape[1], cluster_shape[2]) + m_multiplier = 2 if b_reuse else 1 + tile_descriptions.append( + TileDescription([ + math_inst.instruction_shape[0] * m_multiplier * multiplier_2sm[0], + math_inst.instruction_shape[1] * multiplier_2sm[1], + math_inst.instruction_shape[2] * 2 * multiplier_2sm[2]], + 0, [2, 1, 1], math_inst, min_cc, max_cc, cluster_shape)) + + for cd in cd_data_types: + for layout in layouts: + layout[2][1] = 128 // DataTypeSize[cd["d_type"]] + + # Similar to SM100, for SM107 2SM TMA epilogue with 16-bit output requires CtaN divisible + # by 64 when CtaN > 128. N=160 and N=224 fail the upcast<64> stride-divisibility check + # in the SMEM swizzle layout. Skip tile shapes where CtaN is not 64-aligned. + if DataTypeSize[cd["d_type"]] == 16: + tile_descs_for_dtype = [ + td for td in tile_descriptions if td.threadblock_shape[1] <= 128 or td.threadblock_shape[1] % 64 == 0 + ] + else: + tile_descs_for_dtype = tile_descriptions + if not tile_descs_for_dtype: + continue + + data_type = make_data_type(math_inst, cd) + ops = CreateGemmUniversal3xOperator(manifest, layouts, tile_descs_for_dtype, data_type, + [[kernel_schedule_2sm, epilogue_schedule_2sm]], + tile_schedulers=tile_schedulers, gemm_kind=gemm_kind) + + + +################################################################################################### + def GenerateSM100(manifest, cuda_version): arch_family_cc = ['100f', '101f', '103a', '107f'] if CudaToolkitVersionSatisfies(cuda_version, 13, 0): @@ -12127,6 +12539,12 @@ def GenerateSM100(manifest, cuda_version): GenerateSM100_TensorOp_fp8_UMMA_conv3x(manifest, cuda_version) +def GenerateSM107(manifest, cuda_version): + + GenerateSM107_TensorOp_fp8_UMMA_gemm(manifest, cuda_version) + GenerateSM107_TensorOp_mxf8f6f4_UMMA_gemm_with_block_scaled(manifest, cuda_version) + GenerateSM107_TensorOp_mxnvf4_UMMA_gemm_with_block_scaled(manifest, cuda_version) + def GenerateSM120(manifest, cuda_version): # StreamK is included in regular generation # # @@ -12651,6 +13069,11 @@ def define_parser(): GenerateSM100(manifest, args.cuda_version) GenerateSM120(manifest, args.cuda_version) + rubin_arch_list = ["107a", "107f"] + rubin_enabled_arch = any(arch in rubin_arch_list for arch in archs) + if rubin_enabled_arch: + GenerateSM107(manifest, args.cuda_version) + if 'library' in args.generator_target.split(','): manifest.emit(GeneratorTarget.Library) diff --git a/python/cutlass_library/library.py b/python/cutlass_library/library.py index 9c8891809f..e1f95cc439 100644 --- a/python/cutlass_library/library.py +++ b/python/cutlass_library/library.py @@ -89,8 +89,9 @@ class DataType(enum.Enum): e3m2 = enum_auto() e2m3 = enum_auto() e2m1 = enum_auto() - ue8m0 = enum_auto() - ue4m3 = enum_auto() + ue8m0 = enum_auto() + ue4m3 = enum_auto() + ue5m3 = enum_auto() f16 = enum_auto() bf16 = enum_auto() f32 = enum_auto() @@ -154,8 +155,9 @@ class DataType(enum.Enum): DataType.e2m3: 'e2m3', DataType.e3m2: 'e3m2', DataType.e2m1: 'e2m1', - DataType.ue8m0: 'ue8m0', - DataType.ue4m3: 'ue4m3', + DataType.ue8m0: 'ue8m0', + DataType.ue4m3: 'ue4m3', + DataType.ue5m3: 'ue5m3', DataType.f16: "f16", DataType.bf16: "bf16", DataType.f32: "f32", @@ -203,8 +205,9 @@ class DataType(enum.Enum): DataType.e2m3: 'cutlass::float_e2m3_t', DataType.e3m2: 'cutlass::float_e3m2_t', DataType.e2m1: 'cutlass::float_e2m1_t', - DataType.ue8m0: 'cutlass::float_ue8m0_t', - DataType.ue4m3: 'cutlass::float_ue4m3_t', + DataType.ue8m0: 'cutlass::float_ue8m0_t', + DataType.ue4m3: 'cutlass::float_ue4m3_t', + DataType.ue5m3: 'cutlass::float_ue5m3_t', DataType.f16: "cutlass::half_t", DataType.bf16: "cutlass::bfloat16_t", DataType.f32: "float", @@ -254,6 +257,7 @@ class DataType(enum.Enum): DataType.e2m1: 4, DataType.ue8m0: 8, DataType.ue4m3: 8, + DataType.ue5m3: 8, DataType.f16: 16, DataType.bf16: 16, DataType.f32: 32, @@ -599,6 +603,25 @@ class KernelScheduleType(enum.Enum): PtrArrayMxNvf4UltraTmaWarpSpecialized1SmVs32Sm103TmaPrefetch = enum_auto() PtrArrayMxNvf4UltraTmaWarpSpecialized2SmVs32Sm103TmaPrefetch = enum_auto() + TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithoutBreuse = enum_auto() + TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithoutBreuse = enum_auto() + TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithBreuse = enum_auto() + TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithBreuse = enum_auto() + + TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse = enum_auto() + TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse = enum_auto() + TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse = enum_auto() + TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse = enum_auto() + + TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse = enum_auto() + TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse = enum_auto() + TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse = enum_auto() + TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse = enum_auto() + TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse = enum_auto() + TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse = enum_auto() + TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse = enum_auto() + TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse = enum_auto() + Mxf8f6f4TmaWarpSpecializedCooperativeSm120 = enum_auto() Mxf8f6f4TmaWarpSpecializedPingpongSm120 = enum_auto() Nvf4TmaWarpSpecializedCooperativeSm120 = enum_auto() @@ -722,6 +745,25 @@ class KernelScheduleType(enum.Enum): KernelScheduleType.PtrArrayMxNvf4UltraTmaWarpSpecialized1SmVs32Sm103DisablePrefetch: 'cutlass::gemm::KernelPtrArrayTmaWarpSpecialized1SmBlockScaledMxNvf4UltraVs32Sm103DisablePrefetch', KernelScheduleType.PtrArrayMxNvf4UltraTmaWarpSpecialized2SmVs32Sm103DisablePrefetch: 'cutlass::gemm::KernelPtrArrayTmaWarpSpecialized2SmBlockScaledMxNvf4UltraVs32Sm103DisablePrefetch', + KernelScheduleType.TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithoutBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithoutBreuse', + KernelScheduleType.TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithoutBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithoutBreuse', + KernelScheduleType.TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithBreuse', + KernelScheduleType.TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithBreuse', + + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse', + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse', + + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse', + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse', + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse', + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse: 'cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse', + KernelScheduleType.Mxf8f6f4TmaWarpSpecializedCooperativeSm120: 'cutlass::gemm::KernelTmaWarpSpecializedMxf8f6f4Sm120', KernelScheduleType.Mxf8f6f4TmaWarpSpecializedPingpongSm120: 'cutlass::gemm::KernelTmaWarpSpecializedPingpongMxf8f6f4Sm120', KernelScheduleType.Nvf4TmaWarpSpecializedCooperativeSm120: 'cutlass::gemm::KernelTmaWarpSpecializedNvf4Sm120', @@ -849,6 +891,25 @@ class KernelScheduleType(enum.Enum): KernelScheduleType.PtrArrayMxNvf4UltraTmaWarpSpecialized1SmVs32Sm103TmaPrefetch: '_o_vs32_ultra_1sm_tmapf', KernelScheduleType.PtrArrayMxNvf4UltraTmaWarpSpecialized2SmVs32Sm103TmaPrefetch: '_o_vs32_ultra_2sm_tmapf', + KernelScheduleType.TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithoutBreuse: '_no_breuse_1sm', + KernelScheduleType.TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithoutBreuse: '_no_breuse_2sm', + KernelScheduleType.TmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithBreuse: '_breuse_1sm', + KernelScheduleType.TmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithBreuse: '_breuse_2sm', + + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse: '_no_breuse_1sm', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse: '_no_breuse_2sm', + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse: '_breuse_1sm', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse: '_breuse_2sm', + + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse: '_o_vs16_no_breuse_1sm', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse: '_o_vs16_no_breuse_2sm', + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse: '_o_vs16_breuse_1sm', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse: '_o_vs16_breuse_2sm', + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse: '_o_vs32_no_breuse_1sm', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse: '_o_vs32_no_breuse_2sm', + KernelScheduleType.TmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse: '_o_vs32_breuse_1sm', + KernelScheduleType.TmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse: '_o_vs32_breuse_2sm', + KernelScheduleType.Mxf8f6f4TmaWarpSpecializedCooperativeSm120: '_cooperative_q', KernelScheduleType.Mxf8f6f4TmaWarpSpecializedPingpongSm120: '_pingpong_q', KernelScheduleType.Nvf4TmaWarpSpecializedCooperativeSm120: '_cooperative_o_vs16', @@ -972,12 +1033,12 @@ class EpilogueScheduleType(enum.Enum): class EpilogueFunctor3x(enum.Enum): LinearCombination = enum_auto() - LinearCombinationBlockScaleFactor = enum_auto() + LinearCombinationBlockScaleFactor = enum_auto() # EpilogueFunctor3xTag = { EpilogueFunctor3x.LinearCombination: 'cutlass::epilogue::fusion::LinearCombination', - EpilogueFunctor3x.LinearCombinationBlockScaleFactor: 'cutlass::epilogue::fusion::LinCombBlockScaleFactor', + EpilogueFunctor3x.LinearCombinationBlockScaleFactor: 'cutlass::epilogue::fusion::LinCombBlockScaleFactor', } # TMA epilogues have certain alignment requirements as calculated in get_tma_alignment(data_type) diff --git a/python/cutlass_library/sm107_shapes.py b/python/cutlass_library/sm107_shapes.py new file mode 100644 index 0000000000..345d0ebe43 --- /dev/null +++ b/python/cutlass_library/sm107_shapes.py @@ -0,0 +1,108 @@ +################################################################################################# +# +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. +# +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. +# +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +# +################################################################################################# + + +""" +SM107 UMMA instruction shapes for dense and blockscaled F8/F6/F4 GEMMs. + +Shape format: (M, N, K) + 1SM tiles: M=128, K=64 + 2SM tiles: M=256, K=64 + +Dict values are the minimum tcgen05 instantiation level required to include +a shape (0 = always included at the default level). +""" + +# F8F6F4 dense 1SM: M=128, K=64, N multiples of 16 +SM107_MMA_SHAPES_F8F6F4_DENSE_1SM = { + (128, 16, 64): 4, + (128, 32, 64): 3, + (128, 48, 64): 5, + (128, 64, 64): 2, + (128, 80, 64): 5, + (128, 96, 64): 5, + (128, 112, 64): 5, + (128, 128, 64): 0, + (128, 144, 64): 5, + (128, 160, 64): 5, + (128, 176, 64): 5, + (128, 192, 64): 5, + (128, 208, 64): 5, + (128, 224, 64): 5, + (128, 240, 64): 5, + (128, 256, 64): 0, +} + +# F8F6F4 dense 2SM: M=256, K=64, N multiples of 32 +SM107_MMA_SHAPES_F8F6F4_DENSE_2SM = { + (256, 32, 64): 2, + (256, 64, 64): 2, + (256, 96, 64): 5, + (256, 128, 64): 0, + (256, 160, 64): 5, + (256, 192, 64): 5, + (256, 224, 64): 5, + (256, 256, 64): 0, +} + +# MXF8F6F4 blockscaled 1SM: +SM107_MMA_SHAPES_MXF8F6F4_BLOCKSCALED_1SM = { + (128, 64, 64): 0, + (128, 128, 64): 0, + (128, 192, 64): 0, + (128, 256, 64): 0, +} + +# MXF8F6F4 blockscaled 2SM: +SM107_MMA_SHAPES_MXF8F6F4_BLOCKSCALED_2SM = { + (256, 64, 64): 0, + (256, 128, 64): 0, + (256, 192, 64): 0, + (256, 256, 64): 0, +} + +# MXNVF4 blockscaled 1SM: M=128, K=128 (fp4 doubles the K instruction depth vs f8). Shared by +# both vector sizes (16 and 32) -- they instantiate the same instruction shapes. +SM107_MMA_SHAPES_MXNVF4_BLOCKSCALED_1SM = { + (128, 64, 128): 0, + (128, 128, 128): 0, + (128, 192, 128): 0, + (128, 256, 128): 0, +} + +# MXNVF4 blockscaled 2SM: M=256, K=128. Shared by both vector sizes (16 and 32). +SM107_MMA_SHAPES_MXNVF4_BLOCKSCALED_2SM = { + (256, 64, 128): 0, + (256, 128, 128): 0, + (256, 192, 128): 0, + (256, 256, 128): 0, +} diff --git a/python/cutlass_library/sm107_utils.py b/python/cutlass_library/sm107_utils.py new file mode 100644 index 0000000000..569b047058 --- /dev/null +++ b/python/cutlass_library/sm107_utils.py @@ -0,0 +1,251 @@ +################################################################################################# +# +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. +# +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. +# +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. +# +################################################################################################# + + +""" +Utilities for enumerating CUTLASS library SM107 kernels. +""" + +import warnings +from itertools import product + +try: + from cutlass_library.library import * + from cutlass_library.sm107_shapes import ( + SM107_MMA_SHAPES_F8F6F4_DENSE_1SM, + SM107_MMA_SHAPES_F8F6F4_DENSE_2SM, + SM107_MMA_SHAPES_MXF8F6F4_BLOCKSCALED_1SM, + SM107_MMA_SHAPES_MXF8F6F4_BLOCKSCALED_2SM, + SM107_MMA_SHAPES_MXNVF4_BLOCKSCALED_1SM, + SM107_MMA_SHAPES_MXNVF4_BLOCKSCALED_2SM, + ) + from cutlass_library.sm100_utils import get_tcgen05_level_from_global_level +except ImportError: + from library import * + from sm107_shapes import ( + SM107_MMA_SHAPES_F8F6F4_DENSE_1SM, + SM107_MMA_SHAPES_F8F6F4_DENSE_2SM, + SM107_MMA_SHAPES_MXF8F6F4_BLOCKSCALED_1SM, + SM107_MMA_SHAPES_MXF8F6F4_BLOCKSCALED_2SM, + SM107_MMA_SHAPES_MXNVF4_BLOCKSCALED_1SM, + SM107_MMA_SHAPES_MXNVF4_BLOCKSCALED_2SM, + ) + from sm100_utils import get_tcgen05_level_from_global_level + + +def generate_f8f6f4_math_instructions_sm107( + level: int, + enable_runtime_dtype: bool = True, + enable_compile_time_dtype: bool = False, +): + """ + Generate TensorOp MathInstruction objects for SM107 F8/F6/F4 dense GEMM. + + Args: + level: Global instantiation level. + enable_runtime_dtype: Whether to generate runtime-dtype instructions. + enable_compile_time_dtype: Not yet supported for SM107; emits a warning if True. + + Returns: + (math_instructions_1sm, math_instructions_2sm) + """ + if enable_compile_time_dtype: + warnings.warn( + "SM107 does not support compile-time data types yet. " + "enable_compile_time_dtype=True will be ignored.", + UserWarning, + stacklevel=2, + ) + + tcgen05_level = get_tcgen05_level_from_global_level(level) + + shapes_1sm = [ + shape for shape, min_level in SM107_MMA_SHAPES_F8F6F4_DENSE_1SM.items() + if tcgen05_level >= min_level + ] + shapes_2sm = [ + shape for shape, min_level in SM107_MMA_SHAPES_F8F6F4_DENSE_2SM.items() + if tcgen05_level >= min_level + ] + + math_instructions_1sm = [] + math_instructions_2sm = [] + + for shape in shapes_1sm: + if enable_runtime_dtype: + runtime_types = [ + DataType.f8, + ] + for a_type, b_type in product(runtime_types, repeat=2): + math_instructions_1sm.append( + MathInstruction( + shape, + a_type, b_type, DataType.f32, + OpcodeClass.TensorOp, + MathOperation.multiply_add) + ) + + for shape in shapes_2sm: + if enable_runtime_dtype: + runtime_types = [ + DataType.f8, + ] + for a_type, b_type in product(runtime_types, repeat=2): + math_instructions_2sm.append( + MathInstruction( + shape, + a_type, b_type, DataType.f32, + OpcodeClass.TensorOp, + MathOperation.multiply_add) + ) + + return math_instructions_1sm, math_instructions_2sm + + +def generate_mxf8f6f4_math_instructions_sm107( + level: int, + enable_runtime_dtype: bool = True, + enable_compile_time_dtype: bool = False, +): + """ + Generate BlockScaledTensorOp MathInstruction objects for SM107 MXF8F6F4 blockscaled GEMM. + + Mirrors sm100_utils.generate_mxf8f6f4_math_instructions_sm100(), restricted to the + runtime f8 (a/b dynamic dtype) path only -- SM107 blockscaled currently only supports + a/b = f8 with a runtime dtype, no f6/f4 or compile-time dtypes. + + Args: + level: Global instantiation level. + enable_runtime_dtype: Whether to generate runtime-dtype instructions. + enable_compile_time_dtype: Not yet supported for SM107; emits a warning if True. + + Returns: + (math_instructions_1sm, math_instructions_2sm) + """ + if enable_compile_time_dtype: + warnings.warn( + "SM107 does not support compile-time data types yet. " + "enable_compile_time_dtype=True will be ignored.", + UserWarning, + stacklevel=2, + ) + + tcgen05_level = get_tcgen05_level_from_global_level(level) + + shapes_1sm = [ + shape for shape, min_level in SM107_MMA_SHAPES_MXF8F6F4_BLOCKSCALED_1SM.items() + if tcgen05_level >= min_level + ] + shapes_2sm = [ + shape for shape, min_level in SM107_MMA_SHAPES_MXF8F6F4_BLOCKSCALED_2SM.items() + if tcgen05_level >= min_level + ] + + math_instructions_1sm = [] + math_instructions_2sm = [] + + for shape in shapes_1sm: + if enable_runtime_dtype: + runtime_types = [ + DataType.f8, + ] + for a_type, b_type in product(runtime_types, repeat=2): + math_instructions_1sm.append( + MathInstruction( + shape, + a_type, b_type, DataType.f32, + OpcodeClass.BlockScaledTensorOp, + MathOperation.multiply_add, + DataType.ue8m0) + ) + + for shape in shapes_2sm: + if enable_runtime_dtype: + runtime_types = [ + DataType.f8, + ] + for a_type, b_type in product(runtime_types, repeat=2): + math_instructions_2sm.append( + MathInstruction( + shape, + a_type, b_type, DataType.f32, + OpcodeClass.BlockScaledTensorOp, + MathOperation.multiply_add, + DataType.ue8m0) + ) + + return math_instructions_1sm, math_instructions_2sm + + +def generate_mxnvf4_math_instructions_sm107(level: int): + """ + Generate BlockScaledTensorOp MathInstruction objects for SM107 MXNVF4 blockscaled GEMM. + + a/b are always the compile-time e2m1 type -- SM107 NVF4 has no other 4-bit element + type, so a runtime-dtype (type-erased) a/b is pointless here. Both scale-factor + vector sizes (16 and 32) support all three scale-factor dtypes: ue8m0, ue4m3, ue5m3. + The shapes/instructions themselves don't differ by vector size, so callers reuse the + same (math_instructions_1sm, math_instructions_2sm) lists for both Vs16 and Vs32 + kernel schedules; only the chosen kernel schedule type differs. + + Args: + level: Global instantiation level. + + Returns: + (math_instructions_1sm, math_instructions_2sm) + """ + tcgen05_level = get_tcgen05_level_from_global_level(level) + + sf_types = [DataType.ue8m0, DataType.ue4m3, DataType.ue5m3] + + def _generate(shapes_table): + math_instructions = [] + shapes = [ + shape for shape, min_level in shapes_table.items() + if tcgen05_level >= min_level + ] + for shape in shapes: + for sf_type in sf_types: + math_instructions.append( + MathInstruction( + shape, + DataType.e2m1, DataType.e2m1, DataType.f32, + OpcodeClass.BlockScaledTensorOp, + MathOperation.multiply_add, + sf_type) + ) + return math_instructions + + math_instructions_1sm = _generate(SM107_MMA_SHAPES_MXNVF4_BLOCKSCALED_1SM) + math_instructions_2sm = _generate(SM107_MMA_SHAPES_MXNVF4_BLOCKSCALED_2SM) + + return math_instructions_1sm, math_instructions_2sm diff --git a/test/unit/core/numeric_conversion.cu b/test/unit/core/numeric_conversion.cu index e2d2e71235..1c6c163262 100644 --- a/test/unit/core/numeric_conversion.cu +++ b/test/unit/core/numeric_conversion.cu @@ -48,12 +48,13 @@ namespace kernel { ///////////////////////////////////////////////////////////////////////////////////////////////// /// Simple conversion function -template +template __global__ void convert( cutlass::Array *destination, cutlass::Array const *source) { - cutlass::NumericArrayConverter convert; + cutlass::NumericArrayConverter convert; *destination = convert(*source); } @@ -95,6 +96,8 @@ void run_test(const char dest_name[], const char source_name[], const int range ///////////////////////////////////////////////////////////////////////////////////////////////// +///////////////////////////////////////////////////////////////////////////////////////////////// + template __global__ void convert_with_scale_factor( cutlass::Array *destination, @@ -188,7 +191,9 @@ TEST(NumericConversion, f32x8_to_f16x8_rn) { ///////////////////////////////////////////////////////////////////////////////////////////////// -TEST(NumericConversion, f16_to_f32_rn) { +///////////////////////////////////////////////////////////////////////////////////////////////// + +TEST(NumericConversion, f16_to_f32_rn) { int const kN = 1; using Source = cutlass::half_t; const char source_name[] = "half_t"; diff --git a/test/unit/gemm/device/CMakeLists.txt b/test/unit/gemm/device/CMakeLists.txt index 5b501fadf2..712b030384 100644 --- a/test/unit/gemm/device/CMakeLists.txt +++ b/test/unit/gemm/device/CMakeLists.txt @@ -55,6 +55,10 @@ add_subdirectory(sm100_blockscaled_tensorop_gemm) add_subdirectory(sm100_tensorop_gemm) add_subdirectory(sm100_blockscaled_sparse_tensorop_gemm) add_subdirectory(sm100_sparse_tensorop_gemm) +if (CUTLASS_NVCC_ARCHS MATCHES 107a|107f|109a|109f) + add_subdirectory(sm107_tensorop_gemm) + add_subdirectory(sm107_blockscaled_tensorop_gemm) +endif() add_subdirectory(sm120_blockscaled_sparse_tensorop_gemm) add_subdirectory(sm120_sparse_tensorop_gemm) add_subdirectory(sm120_tensorop_gemm) diff --git a/test/unit/gemm/device/gemm_testbed_3x.hpp b/test/unit/gemm/device/gemm_testbed_3x.hpp index b8d6a3563a..3d87467744 100644 --- a/test/unit/gemm/device/gemm_testbed_3x.hpp +++ b/test/unit/gemm/device/gemm_testbed_3x.hpp @@ -289,7 +289,9 @@ bool initialize_tensor( else if (bits_input <= 8) { if constexpr ( - cute::is_same_v){ + cute::is_same_v || + cute::is_same_v || + cute::is_same_v){ scope_max = 4; scope_min = 1; } @@ -1169,6 +1171,34 @@ struct HostCollectiveMainloop +struct HostCollectiveMainloop, + Gemm, ElementA_, ElementB_> : public + HostCollectiveMainloop, + Gemm, ElementA_, ElementB_> { + using Base = HostCollectiveMainloop, + Gemm, ElementA_, ElementB_>; + HostCollectiveMainloop( + CheckEquality check_relative_equality_ = CheckEquality::EXACT, + cutlass::Distribution::Kind init_A_ = cutlass::Distribution::Uniform, + cutlass::Distribution::Kind init_B_ = cutlass::Distribution::Uniform, + uint64_t seed_ = Base::kDefaultSeed, + typename Base::LayoutTagA::Stride stride_factor_A_ = typename Base::LayoutTagA::Stride(), + typename Base::LayoutTagB::Stride stride_factor_B_ = typename Base::LayoutTagB::Stride() + ) : Base::HostCollectiveMainloop(check_relative_equality_, init_A_, init_B_, seed_, stride_factor_A_, stride_factor_B_) {} +}; + + // // Block Scaled Gemm Input Operands : A , B, scalefactorA, scalefactorB // diff --git a/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/CMakeLists.txt b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/CMakeLists.txt new file mode 100644 index 0000000000..cfb86ca71b --- /dev/null +++ b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/CMakeLists.txt @@ -0,0 +1,57 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. +# +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. +# +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + +if (CUTLASS_NVCC_ARCHS MATCHES 107a) +add_custom_target( + cutlass_test_unit_gemm_device_sm107_bstensorop + DEPENDS + cutlass_test_unit_gemm_device_bstensorop_sm107_mxf8 + cutlass_test_unit_gemm_device_bstensorop_sm107_nvf4 +) + +cutlass_test_unit_gemm_device_add_executable( + cutlass_test_unit_gemm_device_bstensorop_sm107_mxf8 + + BATCH_SOURCES ON + BATCH_SIZE 1 + + mxf8_mxf8_void_f32.cu + mxf8_mxf8_f16_e4m3.cu +) + +cutlass_test_unit_gemm_device_add_executable( + cutlass_test_unit_gemm_device_bstensorop_sm107_nvf4 + + BATCH_SOURCES ON + BATCH_SIZE 1 + + nvf4_nvf4_void_f32.cu + nvf4_nvf4_f16_nvf4.cu +) + +endif() diff --git a/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/mxf8_mxf8_f16_e4m3.cu b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/mxf8_mxf8_f16_e4m3.cu new file mode 100644 index 0000000000..37a060ef8f --- /dev/null +++ b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/mxf8_mxf8_f16_e4m3.cu @@ -0,0 +1,285 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#include + +#include "cutlass/cutlass.h" +#include "cute/tensor.hpp" +#include "cute/atom/mma_atom.hpp" + +#include "cutlass/numeric_types.h" + +#include "cutlass/gemm/device/gemm_universal_adapter.h" +#include "cutlass/gemm/kernel/gemm_universal.hpp" +#include "cutlass/gemm/collective/collective_builder.hpp" + +#include "cutlass/epilogue/dispatch_policy.hpp" +#include "cutlass/epilogue/collective/collective_builder.hpp" + +#include "../../../common/cutlass_unit_test.h" +#include "../gemm_testbed_3x.hpp" + +using namespace cute; + +#if defined(CUTLASS_ARCH_MMA_SM107_SUPPORTED) + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_f16n_e4m3n_bstensorop_f32, 512x256x128_2x2x1_2sm_breuse) { + using ElementA = cutlass::mx_float8_t; + using GmemLayoutA = cutlass::layout::RowMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::mx_float8_t; + using GmemLayoutB = cutlass::layout::ColumnMajor; + constexpr int AlignB = 16; + using ElementC = cutlass::half_t; + using GmemLayoutC = cutlass::layout::ColumnMajor; + constexpr int AlignC = 8; + using ElementD = cutlass::float_e4m3_t; + using GmemLayoutD = cutlass::layout::ColumnMajor; + constexpr int AlignD = 16; + using ElementAccumulator = float; + using ElementCompute = float; + + using MmaTileShape_MNK = Shape<_512,_256,_128>; + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using FusionOperation = cutlass::epilogue::fusion::LinearCombination< + ElementD, ElementCompute, ElementC, ElementCompute + >; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + cutlass::epilogue::TmaWarpSpecialized2Sm, + FusionOperation + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + EXPECT_TRUE(test::gemm::device::TestSmall()); +} + +TEST(SM107_Device_Gemm_ue8m0xe5m2t_ue8m0xe4m3n_bf16t_e4m3t_bstensorop_f32, 256x256x128_2x2x1_1sm_breuse) { + using ElementA = cutlass::mx_float8_t; + using GmemLayoutA = cutlass::layout::RowMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::mx_float8_t; + using GmemLayoutB = cutlass::layout::ColumnMajor; + constexpr int AlignB = 16; + using ElementC = cutlass::bfloat16_t; + using GmemLayoutC = cutlass::layout::RowMajor; + constexpr int AlignC = 8; + using ElementD = cutlass::float_e4m3_t; + using GmemLayoutD = cutlass::layout::RowMajor; + constexpr int AlignD = 16; + using ElementAccumulator = float; + using ElementCompute = float; + + using MmaTileShape_MNK = Shape<_256,_256,_128>; + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using FusionOperation = cutlass::epilogue::fusion::LinearCombination< + ElementD, ElementCompute, ElementC, ElementCompute + >; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + cutlass::epilogue::TmaWarpSpecialized1Sm, + FusionOperation + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + EXPECT_TRUE(test::gemm::device::TestSmall()); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3n_ue8m0xe4m3t_f16t_ue8m0xe4m3t_bstensorop_f32, 256x128x128_2x2x1_1sm_breuse) { + using ElementA = cutlass::mx_float8_t; + using GmemLayoutA = cutlass::layout::ColumnMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::mx_float8_t; + using GmemLayoutB = cutlass::layout::RowMajor; + constexpr int AlignB = 16; + using ElementC = cutlass::half_t; + using GmemLayoutC = cutlass::layout::RowMajor; + constexpr int AlignC = 8; + using ElementD = cutlass::float_e4m3_t; + using GmemLayoutD = cutlass::layout::RowMajor; + constexpr int AlignD = 16; + using ElementAccumulator = float; + using ElementCompute = float; + // Narrow-precision D requires its own scale-factor (SFD) tensor generated by the epilogue. + using ElementSFD = cutlass::float_ue8m0_t; + using GmemLayoutSFD = GmemLayoutD; + constexpr int SFDVectorSize = 32; + + using MmaTileShape_MNK = Shape<_256,_128,_128>; + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor< + SFDVectorSize, + ElementD, ElementCompute, + ElementSFD, GmemLayoutSFD, + ElementC + >; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + cutlass::epilogue::TmaWarpSpecialized1Sm, + FusionOperation + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + EXPECT_TRUE(test::gemm::device::TestSmall()); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe4m3n_bf16t_ue8m0xe4m3t_bstensorop_f32_bias_relu, 512x192x128_2x2x1_2sm_breuse) { + using ElementA = cutlass::mx_float8_t; + using GmemLayoutA = cutlass::layout::RowMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::mx_float8_t; + using GmemLayoutB = cutlass::layout::ColumnMajor; + constexpr int AlignB = 16; + using ElementC = cutlass::bfloat16_t; + using GmemLayoutC = cutlass::layout::RowMajor; + constexpr int AlignC = 8; + using ElementD = cutlass::float_e4m3_t; + using GmemLayoutD = cutlass::layout::RowMajor; + constexpr int AlignD = 16; + using ElementAccumulator = float; + using ElementCompute = float; + // Narrow-precision D requires its own scale-factor (SFD) tensor generated by the epilogue. + using ElementSFD = cutlass::float_ue8m0_t; + using GmemLayoutSFD = GmemLayoutD; + constexpr int SFDVectorSize = 32; + // Bias type + using ElementBias = float; + + using MmaTileShape_MNK = Shape<_512,_192,_128>; + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using FusionOperation = cutlass::epilogue::fusion::LinCombPerColBiasBlockScaleFactor< + SFDVectorSize, ElementD, ElementCompute, + ElementSFD, GmemLayoutSFD, + ElementBias, ElementC + >; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + cutlass::epilogue::TmaWarpSpecialized2Sm, + FusionOperation + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + EXPECT_TRUE(test::gemm::device::TestSmall()); +} + +#endif diff --git a/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/mxf8_mxf8_void_f32.cu b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/mxf8_mxf8_void_f32.cu new file mode 100644 index 0000000000..6a4ff4ff39 --- /dev/null +++ b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/mxf8_mxf8_void_f32.cu @@ -0,0 +1,242 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#include + +#include "cutlass/cutlass.h" +#include "cute/tensor.hpp" +#include "cute/atom/mma_atom.hpp" + +#include "cutlass/numeric_types.h" + +#include "cutlass/gemm/device/gemm_universal_adapter.h" +#include "cutlass/gemm/kernel/gemm_universal.hpp" +#include "cutlass/gemm/collective/collective_builder.hpp" + +#include "cutlass/epilogue/dispatch_policy.hpp" +#include "cutlass/epilogue/collective/collective_builder.hpp" + +#include "../../../common/cutlass_unit_test.h" +#include "../gemm_testbed_3x.hpp" + +using namespace cute; + +#if defined(CUTLASS_ARCH_MMA_SM107_SUPPORTED) + +// Shared structure for SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32: +// A/B are MXFP8 (e4m3), C is unused (void), D is plain f32. Only the mma tile shape, cluster +// shape, and epilogue/mainloop schedules (1SM vs 2SM, with/without B-operand reuse) vary per test. +template < + class MmaTileShape_MNK, + class ClusterShape_MNK, + class EpilogueScheduleType, + class MainloopScheduleType +> +bool run_mxf8_mxf8_void_f32_test() { + using ElementA = cutlass::mx_float8_t; + using GmemLayoutA = cutlass::layout::RowMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::mx_float8_t; + using GmemLayoutB = cutlass::layout::ColumnMajor; + constexpr int AlignB = 16; + using ElementC = void; + using GmemLayoutC = cutlass::layout::ColumnMajor; + constexpr int AlignC = 4; + using ElementD = float; + using GmemLayoutD = cutlass::layout::ColumnMajor; + constexpr int AlignD = 4; + using ElementAccumulator = float; + using ElementCompute = float; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + EpilogueScheduleType + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + MainloopScheduleType + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + + using ProblemShapeType = typename Gemm::GemmKernel::ProblemShape; + ProblemShapeType problem_size{512, 256, 512, /* l */ 1}; + + test::gemm::device::Testbed3x testbed; + return testbed.run(problem_size); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 128x64x128_2x2x1_1sm) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_128,_64,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 256x64x128_2x2x1_2sm) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_256,_64,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 256x64x128_2x2x1_1sm_breuse) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_256,_64,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 512x64x128_2x2x1_2sm_breuse) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_512,_64,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 128x128x128_2x2x1_1sm) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_128,_128,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 256x128x128_2x2x1_2sm) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_256,_128,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 256x128x128_2x2x1_1sm_breuse) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_256,_128,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 512x128x128_2x2x1_2sm_breuse) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_512,_128,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 128x192x128_2x2x1_1sm) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_128,_192,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 256x192x128_2x2x1_2sm) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_256,_192,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 256x192x128_2x2x1_1sm_breuse) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_256,_192,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 512x192x128_2x2x1_2sm_breuse) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_512,_192,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 128x256x128_2x2x1_1sm) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_128,_256,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithoutBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 256x256x128_2x2x1_2sm) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_256,_256,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithoutBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 256x256x128_2x2x1_1sm_breuse) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_256,_256,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxf8f6f4WithBreuse + >())); +} + +TEST(SM107_Device_Gemm_ue8m0xe4m3t_ue8m0xe5m2n_void_f32n_bstensorop_f32, 512x256x128_2x2x1_2sm_breuse) { + EXPECT_TRUE((run_mxf8_mxf8_void_f32_test< + Shape<_512,_256,_128>, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxf8f6f4WithBreuse + >())); +} + +#endif diff --git a/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/nvf4_nvf4_f16_nvf4.cu b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/nvf4_nvf4_f16_nvf4.cu new file mode 100644 index 0000000000..21e3d719d5 --- /dev/null +++ b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/nvf4_nvf4_f16_nvf4.cu @@ -0,0 +1,171 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#include + +#include "cutlass/cutlass.h" +#include "cute/tensor.hpp" +#include "cute/atom/mma_atom.hpp" + +#include "cutlass/numeric_types.h" + +#include "cutlass/gemm/device/gemm_universal_adapter.h" +#include "cutlass/gemm/kernel/gemm_universal.hpp" +#include "cutlass/gemm/collective/collective_builder.hpp" + +#include "cutlass/epilogue/dispatch_policy.hpp" +#include "cutlass/epilogue/collective/collective_builder.hpp" + +#include "../../../common/cutlass_unit_test.h" +#include "../gemm_testbed_3x.hpp" + +using namespace cute; + +#if defined(CUTLASS_ARCH_MMA_SM107_SUPPORTED) + +template < + class MmaTileShape_MNK, + class ClusterShape_MNK, + class EpilogueSchedule, + class MainloopSchedule, + class ElementSFD, + int SFDVectorSize +> +void test_sm107_nvf4_nvf4_f16_nvf4() { + using ElementA = cute::tuple; + using GmemLayoutA = cutlass::layout::RowMajor; + constexpr int AlignA = 32; + using ElementB = cute::tuple; + using GmemLayoutB = cutlass::layout::ColumnMajor; + constexpr int AlignB = 32; + using ElementC = cutlass::half_t; + using GmemLayoutC = cutlass::layout::RowMajor; + constexpr int AlignC = 8; + using ElementD = cutlass::float_e2m1_t; + using GmemLayoutD = cutlass::layout::RowMajor; + constexpr int AlignD = 32; + using ElementAccumulator = float; + using ElementCompute = float; + // Narrow-precision D requires its own scale-factor (SFD) tensor generated by the epilogue. + using GmemLayoutSFD = GmemLayoutD; + + using FusionOperation = cutlass::epilogue::fusion::LinCombBlockScaleFactor< + SFDVectorSize, + ElementD, ElementCompute, + ElementSFD, GmemLayoutSFD, + ElementC + >; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + EpilogueSchedule, + FusionOperation + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + MainloopSchedule + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + + using ProblemShapeType = typename Gemm::GemmKernel::ProblemShape; + ProblemShapeType problem_size{512, 512, 512, /* l */ 1}; + + test::gemm::device::Testbed3x testbed; + EXPECT_TRUE(testbed.run(problem_size)); +} + +// input vs16, without breuse, 1sm, output vs16, SFD = ue8m0 +TEST(SM107_Device_Gemm_ue5m3xe2m1t_ue5m3xe2m1n_f16t_vs16_ue8m0xe2m1t_bstensorop_f32, 512x64x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_f16_nvf4, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse, + cutlass::float_ue8m0_t, 16>(); +} + +// input vs32, with breuse, 1sm, output vs16, SFD = ue4m3 +TEST(SM107_Device_Gemm_ue5m3xe2m1t_ue5m3xe2m1n_f16t_vs16_ue4m3xe2m1t_bstensorop_f32, 256x128x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_f16_nvf4, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse, + cutlass::float_ue4m3_t, 16>(); +} + +// input vs16, without breuse, 2sm, output vs16, SFD = ue5m3 +TEST(SM107_Device_Gemm_ue5m3xe2m1t_ue5m3xe2m1n_f16t_vs16_ue5m3xe2m1t_bstensorop_f32, 256x64x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_f16_nvf4, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse, + cutlass::float_ue5m3_t, 16>(); +} + +// input vs32, without breuse, 1sm, output vs32, SFD = ue8m0 +TEST(SM107_Device_Gemm_ue5m3xe2m1t_ue5m3xe2m1n_f16t_vs32_ue8m0xe2m1t_bstensorop_f32, 128x128x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_f16_nvf4, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse, + cutlass::float_ue8m0_t, 32>(); +} + +// input vs16, without breuse, 2sm, output vs32, SFD = ue4m3 +TEST(SM107_Device_Gemm_ue5m3xe2m1t_ue5m3xe2m1n_f16t_vs32_ue4m3xe2m1t_bstensorop_f32, 256x256x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_f16_nvf4, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse, + cutlass::float_ue4m3_t, 32>(); +} + +// input vs32, with breuse, 2sm, output vs32, SFD = ue5m3 +TEST(SM107_Device_Gemm_ue5m3xe2m1t_ue5m3xe2m1n_f16t_vs32_ue5m3xe2m1t_bstensorop_f32, 512x256x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_f16_nvf4, Shape<_2,_2,_1>, + cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse, + cutlass::float_ue5m3_t, 32>(); +} + +#endif diff --git a/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/nvf4_nvf4_void_f32.cu b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/nvf4_nvf4_void_f32.cu new file mode 100644 index 0000000000..fdb7a84ee8 --- /dev/null +++ b/test/unit/gemm/device/sm107_blockscaled_tensorop_gemm/nvf4_nvf4_void_f32.cu @@ -0,0 +1,280 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#include + +#include "cutlass/cutlass.h" +#include "cute/tensor.hpp" +#include "cute/atom/mma_atom.hpp" + +#include "cutlass/numeric_types.h" + +#include "cutlass/gemm/device/gemm_universal_adapter.h" +#include "cutlass/gemm/kernel/gemm_universal.hpp" +#include "cutlass/gemm/collective/collective_builder.hpp" + +#include "cutlass/epilogue/dispatch_policy.hpp" +#include "cutlass/epilogue/collective/collective_builder.hpp" + +#include "../../../common/cutlass_unit_test.h" +#include "../gemm_testbed_3x.hpp" + +using namespace cute; + +#if defined(CUTLASS_ARCH_MMA_SM107_SUPPORTED) + +template +void test_sm107_nvf4_nvf4_void_f32() { + using ElementA = cute::tuple; + using GmemLayoutA = cutlass::layout::RowMajor; + constexpr int AlignA = 32; + using ElementB = cute::tuple; + using GmemLayoutB = cutlass::layout::ColumnMajor; + constexpr int AlignB = 32; + using ElementC = void; + using GmemLayoutC = cutlass::layout::ColumnMajor; + constexpr int AlignC = 4; + using ElementD = float; + using GmemLayoutD = cutlass::layout::ColumnMajor; + constexpr int AlignD = 4; + using ElementAccumulator = float; + using ElementCompute = float; + + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + EpilogueSchedule + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassBlockScaledTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + MainloopSchedule + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + + using ProblemShapeType = typename Gemm::GemmKernel::ProblemShape; + ProblemShapeType problem_size{512, 256, 512, /* l */ 1}; + + test::gemm::device::Testbed3x testbed; + EXPECT_TRUE(testbed.run(problem_size)); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 128x64x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue4m3xe2m1t_ue4m3xe2m1n_void_f32n_bstensorop_f32_vs16, 128x64x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse, + cutlass::float_ue4m3_t>(); +} + +TEST(SM107_Device_Gemm_ue5m3xe2m1t_ue5m3xe2m1n_void_f32n_bstensorop_f32_vs16, 128x64x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse, + cutlass::float_ue5m3_t>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 128x128x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 128x192x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 128x256x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 256x64x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 256x128x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 256x192x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 256x256x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 256x64x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 256x128x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 256x192x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 256x256x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs16WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 512x64x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 512x128x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 512x192x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs16, 512x256x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs16WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 128x64x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 128x128x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 128x192x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 128x256x256_2x2x1_1sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 256x64x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 256x128x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 256x192x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 256x256x256_2x2x1_2sm) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithoutBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 256x64x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 256x128x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 256x192x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 256x256x256_2x2x1_1sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized1Sm, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107BlockScaledMxNvf4Vs32WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 512x64x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 512x128x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 512x192x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse>(); +} + +TEST(SM107_Device_Gemm_ue8m0xe2m1t_ue8m0xe2m1n_void_f32n_bstensorop_f32_vs32, 512x256x256_2x2x1_2sm_breuse) { + test_sm107_nvf4_nvf4_void_f32, cutlass::epilogue::TmaWarpSpecialized2Sm, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107BlockScaledMxNvf4Vs32WithBreuse>(); +} + +#endif diff --git a/test/unit/gemm/device/sm107_tensorop_gemm/CMakeLists.txt b/test/unit/gemm/device/sm107_tensorop_gemm/CMakeLists.txt new file mode 100644 index 0000000000..ab134b6e09 --- /dev/null +++ b/test/unit/gemm/device/sm107_tensorop_gemm/CMakeLists.txt @@ -0,0 +1,45 @@ +# Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: BSD-3-Clause +# +# Redistribution and use in source and binary forms, with or without +# modification, are permitted provided that the following conditions are met: +# +# 1. Redistributions of source code must retain the above copyright notice, this +# list of conditions and the following disclaimer. +# +# 2. Redistributions in binary form must reproduce the above copyright notice, +# this list of conditions and the following disclaimer in the documentation +# and/or other materials provided with the distribution. +# +# 3. Neither the name of the copyright holder nor the names of its +# contributors may be used to endorse or promote products derived from +# this software without specific prior written permission. +# +# THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" +# AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE +# IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE +# DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE +# FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL +# DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR +# SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER +# CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, +# OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE +# OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + + +if (CUTLASS_NVCC_ARCHS MATCHES 107a) +add_custom_target( + cutlass_test_unit_gemm_device_sm107_tensorop + DEPENDS + cutlass_test_unit_gemm_device_tensorop_sm107_fp8 +) + +cutlass_test_unit_gemm_device_add_executable( + cutlass_test_unit_gemm_device_tensorop_sm107_fp8 + + BATCH_SOURCES ON + BATCH_SIZE 1 + + f8_f8_void_f32.cu +) +endif() diff --git a/test/unit/gemm/device/sm107_tensorop_gemm/f8_f8_void_f32.cu b/test/unit/gemm/device/sm107_tensorop_gemm/f8_f8_void_f32.cu new file mode 100644 index 0000000000..70ebced1fc --- /dev/null +++ b/test/unit/gemm/device/sm107_tensorop_gemm/f8_f8_void_f32.cu @@ -0,0 +1,250 @@ +/*************************************************************************************************** + * Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#include + +#include "cutlass/cutlass.h" +#include "cute/tensor.hpp" +#include "cute/atom/mma_atom.hpp" + +#include "cutlass/numeric_types.h" + +#include "cutlass/gemm/device/gemm_universal_adapter.h" +#include "cutlass/gemm/kernel/gemm_universal.hpp" +#include "cutlass/gemm/collective/collective_builder.hpp" + +#include "cutlass/epilogue/dispatch_policy.hpp" +#include "cutlass/epilogue/collective/collective_builder.hpp" + +#include "../../../common/cutlass_unit_test.h" +#include "../gemm_testbed_3x.hpp" + +using namespace cute; + +#if defined(CUTLASS_ARCH_MMA_SM107_SUPPORTED) + +TEST(SM107_Device_Gemm_e4m3t_e4m3n_void_f32n_tensor_op_f32, 128x128x128_2x2x1_1sm) { + using ElementA = cutlass::float_e4m3_t; + using GmemLayoutA = cutlass::layout::RowMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::float_e4m3_t; + using GmemLayoutB = cutlass::layout::ColumnMajor; + constexpr int AlignB = 16; + using ElementC = void; + using GmemLayoutC = cutlass::layout::ColumnMajor; + constexpr int AlignC = 4; + using ElementD = float; + using GmemLayoutD = cutlass::layout::ColumnMajor; + constexpr int AlignD = 4; + using ElementAccumulator = float; + using ElementCompute = float; + + using MmaTileShape_MNK = Shape<_128,_128,_128>; + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + cutlass::epilogue::TmaWarpSpecialized1Sm + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithoutBreuse + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + EXPECT_TRUE(test::gemm::device::TestSmall()); +} + +TEST(SM107_Device_Gemm_e4m3t_e4m3n_void_f32n_tensor_op_f32, 128x128x128_2x2x1_2sm) { + using ElementA = cutlass::float_e4m3_t; + using GmemLayoutA = cutlass::layout::ColumnMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::float_e4m3_t; + using GmemLayoutB = cutlass::layout::ColumnMajor; + constexpr int AlignB = 16; + using ElementC = void; + using GmemLayoutC = cutlass::layout::ColumnMajor; + constexpr int AlignC = 4; + using ElementD = float; + using GmemLayoutD = cutlass::layout::ColumnMajor; + constexpr int AlignD = 4; + using ElementAccumulator = float; + using ElementCompute = float; + + using MmaTileShape_MNK = Shape<_256,_128,_128>; + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + cutlass::epilogue::TmaWarpSpecialized2Sm + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithoutBreuse + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + EXPECT_TRUE(test::gemm::device::TestSmall()); +} + +TEST(SM107_Device_Gemm_e4m3t_e4m3n_void_f32n_tensor_op_f32, 256x128x128_2x2x1_1sm_breuse) { + using ElementA = cutlass::float_e4m3_t; + using GmemLayoutA = cutlass::layout::RowMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::float_e4m3_t; + using GmemLayoutB = cutlass::layout::RowMajor; + constexpr int AlignB = 16; + using ElementC = void; + using GmemLayoutC = cutlass::layout::ColumnMajor; + constexpr int AlignC = 4; + using ElementD = float; + using GmemLayoutD = cutlass::layout::ColumnMajor; + constexpr int AlignD = 4; + using ElementAccumulator = float; + using ElementCompute = float; + + using MmaTileShape_MNK = Shape<_256,_128,_128>; + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + cutlass::epilogue::TmaWarpSpecialized1Sm + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelTmaWarpSpecialized1SmSm107DenseGemmf8f6f4WithBreuse + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + EXPECT_TRUE(test::gemm::device::TestSmall()); +} + +TEST(SM107_Device_Gemm_e4m3t_e4m3n_void_f32n_tensor_op_f32, 512x128x128_2x2x1_2sm_breuse) { + using ElementA = cutlass::float_e4m3_t; + using GmemLayoutA = cutlass::layout::ColumnMajor; + constexpr int AlignA = 16; + using ElementB = cutlass::float_e4m3_t; + using GmemLayoutB = cutlass::layout::RowMajor; + constexpr int AlignB = 16; + using ElementC = void; + using GmemLayoutC = cutlass::layout::ColumnMajor; + constexpr int AlignC = 4; + using ElementD = float; + using GmemLayoutD = cutlass::layout::ColumnMajor; + constexpr int AlignD = 4; + using ElementAccumulator = float; + using ElementCompute = float; + + using MmaTileShape_MNK = Shape<_512,_128,_128>; + using ClusterShape_MNK = Shape<_2,_2,_1>; + + using CollectiveEpilogue = typename cutlass::epilogue::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassTensorOp, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::epilogue::collective::EpilogueTileAuto, + ElementAccumulator, ElementCompute, + ElementC, GmemLayoutC, AlignC, + ElementD, GmemLayoutD, AlignD, + cutlass::epilogue::TmaWarpSpecialized2Sm + >::CollectiveOp; + + using CollectiveMainloop = typename cutlass::gemm::collective::CollectiveBuilder< + cutlass::arch::Sm107, cutlass::arch::OpClassTensorOp, + ElementA, GmemLayoutA, AlignA, + ElementB, GmemLayoutB, AlignB, + ElementAccumulator, + MmaTileShape_MNK, ClusterShape_MNK, + cutlass::gemm::collective::StageCountAutoCarveout(sizeof(typename CollectiveEpilogue::SharedStorage))>, + cutlass::gemm::KernelTmaWarpSpecialized2SmSm107DenseGemmf8f6f4WithBreuse + >::CollectiveOp; + + using GemmKernel = cutlass::gemm::kernel::GemmUniversal< + Shape, + CollectiveMainloop, + CollectiveEpilogue + >; + + using Gemm = cutlass::gemm::device::GemmUniversalAdapter; + EXPECT_TRUE(test::gemm::device::TestSmall()); +} + +#endif // CUTLASS_ARCH_MMA_SM107_SUPPORTED diff --git a/test/unit/transform/device/CMakeLists.txt b/test/unit/transform/device/CMakeLists.txt index 4735f5837b..a6f602375c 100644 --- a/test/unit/transform/device/CMakeLists.txt +++ b/test/unit/transform/device/CMakeLists.txt @@ -63,8 +63,8 @@ add_custom_target( cutlass_test_unit_sm100_structured_sparse_gemm_compressor_f16 cutlass_test_unit_sm100_structured_sparse_gemm_compressor_f8 cutlass_test_unit_sm100_structured_sparse_gemm_compressor_f6 - cutlass_test_unit_sm100_structured_sparse_gemm_compressor_f4_qmma - cutlass_test_unit_sm100_structured_sparse_gemm_compressor_f4_omma + cutlass_test_unit_sm100_structured_sparse_gemm_compressor_mxf8f6f4 + cutlass_test_unit_sm100_structured_sparse_gemm_compressor_mxf4 ) cutlass_test_unit_add_executable( @@ -92,14 +92,26 @@ cutlass_test_unit_add_executable( ) cutlass_test_unit_add_executable( - cutlass_test_unit_sm100_structured_sparse_gemm_compressor_f4_qmma + cutlass_test_unit_sm100_structured_sparse_gemm_compressor_mxf8f6f4 - sm100_sparse_gemm_compressor_f4_qmma.cu + sm100_sparse_gemm_compressor_mxf8f6f4.cu ) cutlass_test_unit_add_executable( - cutlass_test_unit_sm100_structured_sparse_gemm_compressor_f4_omma + cutlass_test_unit_sm100_structured_sparse_gemm_compressor_mxf4 - sm100_sparse_gemm_compressor_f4_omma.cu + sm100_sparse_gemm_compressor_mxf4.cu +) + +add_custom_target( + cutlass_test_unit_sm107_structured_sparse_gemm_compressor + DEPENDS + cutlass_test_unit_sm107_structured_sparse_gemm_compressor_k96_mxf4 +) + +cutlass_test_unit_add_executable( + cutlass_test_unit_sm107_structured_sparse_gemm_compressor_k96_mxf4 + + sm107_sparse_gemm_compressor_k96_mxf4.cu ) diff --git a/test/unit/transform/device/sm100_sparse_gemm_compressor_f4_omma.cu b/test/unit/transform/device/sm100_sparse_gemm_compressor_mxf4.cu similarity index 96% rename from test/unit/transform/device/sm100_sparse_gemm_compressor_f4_omma.cu rename to test/unit/transform/device/sm100_sparse_gemm_compressor_mxf4.cu index 24393dfc38..9132b9463a 100644 --- a/test/unit/transform/device/sm100_sparse_gemm_compressor_f4_omma.cu +++ b/test/unit/transform/device/sm100_sparse_gemm_compressor_mxf4.cu @@ -38,14 +38,14 @@ /////////////////////////////////////////////////////////////////////////////////////////////////// // * Test Plan -// ElementA : fp4 (qmma), fp4 (omma) +// ElementA : mxf8f6f4, mxf4 // LayoutA : row / col // Gemm : 1x 2x 3x multiplier of alignment requirement. corner case that smaller than alignment requirement /////////////////////////////////////////////////////////////////////////////////////////////////// #if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED) -TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, omma_f4_t) +TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, mxf4_t) { // Test Settings using ElementA = cutlass::float_e2m1_t; @@ -67,7 +67,7 @@ TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, omma_f4_t) EXPECT_TRUE(testbed.run_auto()); } -TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, omma_f4_runtimedtype_t) +TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, mxf4_runtimedtype_t) { // Test Settings using ElementA = cutlass::type_erased_dynamic_float4_t; diff --git a/test/unit/transform/device/sm100_sparse_gemm_compressor_f4_qmma.cu b/test/unit/transform/device/sm100_sparse_gemm_compressor_mxf8f6f4.cu similarity index 99% rename from test/unit/transform/device/sm100_sparse_gemm_compressor_f4_qmma.cu rename to test/unit/transform/device/sm100_sparse_gemm_compressor_mxf8f6f4.cu index 954b09a1d5..8050658aec 100644 --- a/test/unit/transform/device/sm100_sparse_gemm_compressor_f4_qmma.cu +++ b/test/unit/transform/device/sm100_sparse_gemm_compressor_mxf8f6f4.cu @@ -38,7 +38,7 @@ /////////////////////////////////////////////////////////////////////////////////////////////////// // * Test Plan -// ElementA : fp4 (qmma), fp4 (omma) +// ElementA : mxf8f6f4, mxf4 // LayoutA : row / col // Gemm : 1x 2x 3x multiplier of alignment requirement. corner case that smaller than alignment requirement /////////////////////////////////////////////////////////////////////////////////////////////////// diff --git a/test/unit/transform/device/sm107_sparse_gemm_compressor_k96_mxf4.cu b/test/unit/transform/device/sm107_sparse_gemm_compressor_k96_mxf4.cu new file mode 100644 index 0000000000..e3a4cdf1b9 --- /dev/null +++ b/test/unit/transform/device/sm107_sparse_gemm_compressor_k96_mxf4.cu @@ -0,0 +1,153 @@ +/*************************************************************************************************** + * Copyright (c) 2024 - 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: BSD-3-Clause + * + * Redistribution and use in source and binary forms, with or without + * modification, are permitted provided that the following conditions are met: + * + * 1. Redistributions of source code must retain the above copyright notice, this + * list of conditions and the following disclaimer. + * + * 2. Redistributions in binary form must reproduce the above copyright notice, + * this list of conditions and the following disclaimer in the documentation + * and/or other materials provided with the distribution. + * + * 3. Neither the name of the copyright holder nor the names of its + * contributors may be used to endorse or promote products derived from + * this software without specific prior written permission. + * + * THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" + * AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE + * IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE + * DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE + * FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL + * DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR + * SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER + * CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, + * OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE + * OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. + * + **************************************************************************************************/ + +#include "cute/arch/mma_sm100_desc.hpp" // cute::UMMA::Major +#include "cutlass/gemm/collective/builders/sm100_common.inl" // tag_to_umma_major_A +#include "cutlass/gemm/collective/builders/sm107_sparse_config.inl" // Sm107GemmSparseConfig +#include "cutlass/gemm/collective/builders/sm1xx_sparse_config.inl" // Sm107GemmSparseConfig +#include "cutlass/transform/kernel/sparse_gemm_compressor.hpp" // StructuredSparseCompressor +#include "cutlass/transform/device/transform_universal_adapter.hpp" // TransformUniversalAdapter +#include "testbed_sparse_gemm_compressor.hpp" // TestbedSparseGemmCompressor + +/////////////////////////////////////////////////////////////////////////////////////////////////// +// * Test Plan +// ElementA : mxf4 +// LayoutA : row (mxf4 dont support transposeA) +// Gemm : 1x 2x 3x multiplier of alignment requirement. corner case that smaller than alignment requirement +/////////////////////////////////////////////////////////////////////////////////////////////////// + +#if defined(CUTLASS_ARCH_MMA_SM100_SUPPORTED) || defined(CUTLASS_ARCH_MMA_SM107_SUPPORTED) + +TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, k96_mxf4_t_128x32) +{ + // Test Settings + using ElementA = cutlass::float_e2m1_t; + using LayoutATag = cutlass::layout::RowMajor; + + // Deduct From Test Setting + using ElementAMma = cute::sparse_elem<2, ElementA>; + using ElementEMma = cute::sparse_elem<8, uint8_t>; + + using Sm107SparseConfig = cutlass::Sm107GemmSparseConfig; + + using CompressorKernel = cutlass::transform::kernel:: + StructuredSparseCompressor, + ElementA, + LayoutATag, + Sm107SparseConfig, + cutlass::arch::Sm100>; + + using Compressor = cutlass::transform::device::TransformUniversalAdapter; + + // Test Bed + test::transform::device::TestbedSparseGemmCompressor testbed; + EXPECT_TRUE(testbed.run({128, 128, 32, 1})); +} + +TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, k96_mxf4_t_128x64) +{ + // Test Settings + using ElementA = cutlass::float_e2m1_t; + using LayoutATag = cutlass::layout::RowMajor; + + // Deduct From Test Setting + using ElementAMma = cute::sparse_elem<2, ElementA>; + using ElementEMma = cute::sparse_elem<8, uint8_t>; + + using Sm107SparseConfig = cutlass::Sm107GemmSparseConfig; + + using CompressorKernel = cutlass::transform::kernel:: + StructuredSparseCompressor, + ElementA, + LayoutATag, + Sm107SparseConfig, + cutlass::arch::Sm100>; + + using Compressor = cutlass::transform::device::TransformUniversalAdapter; + + // Test Bed + test::transform::device::TestbedSparseGemmCompressor testbed; + EXPECT_TRUE(testbed.run({128, 128, 64, 1})); +} + +TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, k96_mxf4_t_128x128) +{ + // Test Settings + using ElementA = cutlass::float_e2m1_t; + using LayoutATag = cutlass::layout::RowMajor; + + // Deduct From Test Setting + using ElementAMma = cute::sparse_elem<2, ElementA>; + using ElementEMma = cute::sparse_elem<8, uint8_t>; + + using Sm107SparseConfig = cutlass::Sm107GemmSparseConfig; + + using CompressorKernel = cutlass::transform::kernel:: + StructuredSparseCompressor, + ElementA, + LayoutATag, + Sm107SparseConfig, + cutlass::arch::Sm100>; + + using Compressor = cutlass::transform::device::TransformUniversalAdapter; + + // Test Bed + test::transform::device::TestbedSparseGemmCompressor testbed; + EXPECT_TRUE(testbed.run({128, 128, 128, 1})); +} + +TEST(SM100_Structured_Sparse_Gemm_Compressor_Device, k96_mxf4_t) +{ + // Test Settings + using ElementA = cutlass::float_e2m1_t; + using LayoutATag = cutlass::layout::RowMajor; + + // Deduct From Test Setting + using ElementAMma = cute::sparse_elem<2, ElementA>; + using ElementEMma = cute::sparse_elem<8, uint8_t>; + + using Sm107SparseConfig = cutlass::Sm107GemmSparseConfig; + + using CompressorKernel = cutlass::transform::kernel:: + StructuredSparseCompressor, + ElementA, + LayoutATag, + Sm107SparseConfig, + cutlass::arch::Sm100>; + + using Compressor = cutlass::transform::device::TransformUniversalAdapter; + + // Test Bed + test::transform::device::TestbedSparseGemmCompressor testbed; + EXPECT_TRUE(testbed.run_auto()); +} + +#endif diff --git a/tools/library/include/cutlass/library/types.h b/tools/library/include/cutlass/library/types.h index 0c6c9e1ca3..d190e30a37 100644 --- a/tools/library/include/cutlass/library/types.h +++ b/tools/library/include/cutlass/library/types.h @@ -88,8 +88,9 @@ enum class NumericTypeID { kFE2M3, kFE3M2, kFE2M1, - kFUE8M0, - kFUE4M3, + kFUE8M0, + kFUE4M3, + kFUE5M3, kF8, kF6, kF4, diff --git a/tools/library/src/library_internal.h b/tools/library/src/library_internal.h index 84f2bbe70c..407d750d9f 100644 --- a/tools/library/src/library_internal.h +++ b/tools/library/src/library_internal.h @@ -136,6 +136,10 @@ template <> struct NumericTypeMap { static NumericTypeID const kId = NumericTypeID::kFUE4M3; }; +template <> struct NumericTypeMap { + static NumericTypeID const kId = NumericTypeID::kFUE5M3; +}; + template <> struct NumericTypeMap { static NumericTypeID const kId = NumericTypeID::kU16; diff --git a/tools/library/src/util.cu b/tools/library/src/util.cu index 1277c3d084..272bbc751b 100644 --- a/tools/library/src/util.cu +++ b/tools/library/src/util.cu @@ -505,6 +505,7 @@ NumericTypeID_enumerants[] = { {"fe2m1", "FE2M1", NumericTypeID::kFE2M1}, {"fue8m0", "FUE8M0", NumericTypeID::kFUE8M0}, {"fue4m3", "FUE4M3", NumericTypeID::kFUE4M3}, + {"fue5m3", "FUE5M3", NumericTypeID::kFUE5M3}, {"f16", "F16", NumericTypeID::kF16}, {"bf16", "BF16", NumericTypeID::kBF16}, {"f32", "F32", NumericTypeID::kF32}, @@ -577,6 +578,7 @@ int sizeof_bits(NumericTypeID type) { case NumericTypeID::kFE2M1: return 4; case NumericTypeID::kFUE8M0: return 8; case NumericTypeID::kFUE4M3: return 8; + case NumericTypeID::kFUE5M3: return 8; case NumericTypeID::kF16: return 16; case NumericTypeID::kBF16: return 16; case NumericTypeID::kTF32: return 32; @@ -665,6 +667,7 @@ bool is_signed_type(NumericTypeID type) { case NumericTypeID::kFE2M1: return true; case NumericTypeID::kFUE8M0: return false; case NumericTypeID::kFUE4M3: return false; + case NumericTypeID::kFUE5M3: return false; case NumericTypeID::kF16: return true; case NumericTypeID::kBF16: return true; case NumericTypeID::kTF32: return true; @@ -705,6 +708,7 @@ bool is_float_type(NumericTypeID type) { case NumericTypeID::kFE2M1: return true; case NumericTypeID::kFUE8M0: return true; case NumericTypeID::kFUE4M3: return true; + case NumericTypeID::kFUE5M3: return true; case NumericTypeID::kF16: return true; case NumericTypeID::kBF16: return true; case NumericTypeID::kTF32: return true; @@ -1288,6 +1292,13 @@ bool lexical_cast(std::vector &bytes, NumericTypeID type, std::string c *reinterpret_cast(bytes.data()) = static_cast(tmp); } break; + case NumericTypeID::kFUE5M3: + { + float tmp; + ss >> tmp; + *reinterpret_cast(bytes.data()) = static_cast(tmp); + } + break; case NumericTypeID::kF16: { float tmp; @@ -1468,6 +1479,12 @@ std::string lexical_cast(std::vector &bytes, NumericTypeID type) { ss << tmp; } break; + case NumericTypeID::kFUE5M3: + { + float tmp = *reinterpret_cast(bytes.data()); + ss << tmp; + } + break; case NumericTypeID::kF16: { float tmp = *reinterpret_cast(bytes.data()); @@ -1646,6 +1663,11 @@ bool cast_from_int64(std::vector &bytes, NumericTypeID type, int64_t sr *reinterpret_cast(bytes.data()) = static_cast(float(src)); } break; + case NumericTypeID::kFUE5M3: + { + *reinterpret_cast(bytes.data()) = static_cast(float(src)); + } + break; case NumericTypeID::kF16: { *reinterpret_cast(bytes.data()) = static_cast(float(src)); @@ -1782,6 +1804,11 @@ bool cast_from_uint64(std::vector &bytes, NumericTypeID type, uint64_t *reinterpret_cast(bytes.data()) = static_cast(float(src)); } break; + case NumericTypeID::kFUE5M3: + { + *reinterpret_cast(bytes.data()) = static_cast(float(src)); + } + break; case NumericTypeID::kF16: { *reinterpret_cast(bytes.data()) = static_cast(float(src)); @@ -1919,6 +1946,11 @@ bool cast_from_double(std::vector &bytes, NumericTypeID type, double sr *reinterpret_cast(bytes.data()) = static_cast(float(src)); } break; + case NumericTypeID::kFUE5M3: + { + *reinterpret_cast(bytes.data()) = static_cast(float(src)); + } + break; case NumericTypeID::kF16: { *reinterpret_cast(bytes.data()) = static_cast(float(src)); diff --git a/tools/profiler/src/device_allocation.cu b/tools/profiler/src/device_allocation.cu index 5cea42e116..c1beb76689 100644 --- a/tools/profiler/src/device_allocation.cu +++ b/tools/profiler/src/device_allocation.cu @@ -655,6 +655,14 @@ void DeviceAllocation::initialize_random_device(int seed, Distribution dist) { dist ); break; + case library::NumericTypeID::kFUE5M3: + cutlass::reference::device::BlockFillRandom( + reinterpret_cast(pointer_), + capacity_, + seed, + dist + ); + break; case library::NumericTypeID::kFUE8M0: cutlass::reference::device::BlockFillRandom( reinterpret_cast(pointer_), @@ -851,6 +859,14 @@ void DeviceAllocation::initialize_random_host(int seed, Distribution dist) { dist ); break; + case library::NumericTypeID::kFUE5M3: + cutlass::reference::host::BlockFillRandom( + reinterpret_cast(host_data.data()), + capacity_, + seed, + dist + ); + break; case library::NumericTypeID::kFE2M3: cutlass::reference::host::BlockFillRandom( @@ -1112,6 +1128,14 @@ void DeviceAllocation::initialize_sequential_device(Distribution dist) { static_cast(dist.sequential.start) ); break; + case library::NumericTypeID::kFUE5M3: + cutlass::reference::device::BlockFillSequential( + reinterpret_cast(pointer_), + capacity_, + static_cast(dist.sequential.delta), + static_cast(dist.sequential.start) + ); + break; case library::NumericTypeID::kFE2M3: cutlass::reference::device::BlockFillSequential( @@ -1384,6 +1408,14 @@ void DeviceAllocation::initialize_sequential_host(Distribution dist) { static_cast(dist.sequential.start) ); break; + case library::NumericTypeID::kFUE5M3: + cutlass::reference::host::BlockFillSequential( + reinterpret_cast(host_data.data()), + capacity_, + static_cast(dist.sequential.delta), + static_cast(dist.sequential.start) + ); + break; case library::NumericTypeID::kFE2M3: cutlass::reference::host::BlockFillSequential( @@ -1718,6 +1750,11 @@ bool DeviceAllocation::block_compare_equal( reinterpret_cast(ptr_A), reinterpret_cast(ptr_B), capacity); + case library::NumericTypeID::kFUE5M3: + return reference::device::BlockCompareEqual( + reinterpret_cast(ptr_A), + reinterpret_cast(ptr_B), + capacity); case library::NumericTypeID::kFUE8M0: return reference::device::BlockCompareEqual( reinterpret_cast(ptr_A), @@ -1914,6 +1951,13 @@ bool DeviceAllocation::block_compare_relatively_equal( capacity, static_cast(epsilon), static_cast(nonzero_floor)); + case library::NumericTypeID::kFUE5M3: + return reference::device::BlockCompareRelativelyEqual( + reinterpret_cast(ptr_A), + reinterpret_cast(ptr_B), + capacity, + static_cast(epsilon), + static_cast(nonzero_floor)); case library::NumericTypeID::kFUE8M0: return reference::device::BlockCompareRelativelyEqual( reinterpret_cast(ptr_A), @@ -2291,6 +2335,9 @@ void DeviceAllocation::write_tensor_csv( case library::NumericTypeID::kFUE4M3: write_tensor_csv_static_type(out, *this); break; + case library::NumericTypeID::kFUE5M3: + write_tensor_csv_static_type(out, *this); + break; case library::NumericTypeID::kFE2M3: write_tensor_csv_static_type(out, *this); @@ -2477,6 +2524,9 @@ void DeviceAllocation::fill_device(double val = 0.0) { case library::NumericTypeID::kFUE4M3: tensor_fill(*this, static_cast(val)); break; + case library::NumericTypeID::kFUE5M3: + tensor_fill(*this, static_cast(val)); + break; case library::NumericTypeID::kFUE8M0: tensor_fill(*this, static_cast(val)); @@ -2596,6 +2646,13 @@ void DeviceAllocation::fill_host(double val = 0.0) { static_cast(val) ); break; + case library::NumericTypeID::kFUE5M3: + cutlass::reference::host::BlockFill( + reinterpret_cast(host_data.data()), + capacity_, + static_cast(val) + ); + break; case library::NumericTypeID::kFUE8M0: cutlass::reference::host::BlockFill( diff --git a/tools/profiler/src/device_context.cu b/tools/profiler/src/device_context.cu index 760a2591f7..6041a38b2a 100644 --- a/tools/profiler/src/device_context.cu +++ b/tools/profiler/src/device_context.cu @@ -171,7 +171,11 @@ DeviceAllocation *DeviceContext::allocate_and_initialize_tensor( case library::NumericTypeID::kFUE4M3: data_distribution.set_uniform(1, 4, 0); break; - + + case library::NumericTypeID::kFUE5M3: + data_distribution.set_uniform(1, 4, 0); + break; + case library::NumericTypeID::kF16: data_distribution.set_uniform(-3, 3, 0); break;