[Backend] Support TMA lowering for arbitrary (swizzled) SMEM layout - #2380
Conversation
|
👋 Hi! Thank you for contributing to the TileLang project. Please remember to run We appreciate you taking this step! Our team will review your contribution, and we look forward to your awesome work! 🚀 |
|
No actionable comments were generated in the recent review. 🎉 ℹ️ Recent review info⚙️ Run configurationConfiguration used: defaults Review profile: CHILL Plan: Pro Run ID: 📒 Files selected for processing (11)
✅ Files skipped from review due to trivial changes (1)
🚧 Files skipped from review as they are similar to previous changes (7)
📝 WalkthroughWalkthroughIntroduces a full CuTe-style layout IR (Swizzle, IntTuple, Layout, ComposedLayout) in C++ with symbolic recovery from TileLang layouts via GF(2) sampling and parity-based proof machinery. Wires the IR to Python via TVM FFI. Refactors TMA bulk-copy descriptor construction and IsTmaCompatibleLayout to consume the new IR, and rewrites MMA swizzle index generation to return three-component tuples with XOR-based swizzle arithmetic. ChangesCuTe Layout IR, TMA Swizzle Recovery, and Consumers
Sequence Diagram(s)sequenceDiagram
participant Kernel as TileLang Kernel
participant LowerBulk as Copy::LowerBulk
participant ComposedLayoutFromTileLang
participant CuTeAlgebra as CuTe Coalesce/Compose/RightInverse
participant TMADescriptor as CU_TENSOR_MAP descriptor
participant GPU as CUDA TMA Unit
Kernel->>LowerBulk: annotated shared-mem buffer + global tensor
LowerBulk->>ComposedLayoutFromTileLang: shared TileLang Layout
ComposedLayoutFromTileLang-->>LowerBulk: ComposedLayout (b_bits, m_base, s_shift, plain layout)
LowerBulk->>LowerBulk: map swizzle → CU_TENSOR_MAP_SWIZZLE_*
LowerBulk->>CuTeAlgebra: compose tile→smem and tile→global, RightInverse, Coalesce(max_extent)
CuTeAlgebra-->>LowerBulk: TMA box dims, global strides, mode routing tables
LowerBulk->>TMADescriptor: set box, global_shape, global_stride, swizzle mode
TMADescriptor-->>LowerBulk: descriptor handle
LowerBulk->>LowerBulk: build make_shared_offset / make_tma_coords lambdas
LowerBulk->>GPU: emit tma_load(descriptor, smem_addr, tma_coords) in rest_size loop
Estimated code review effort🎯 5 (Critical) | ⏱️ ~120 minutes Suggested reviewers
🚥 Pre-merge checks | ✅ 4 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (4 passed)
✏️ Tip: You can configure your own custom pre-merge checks in the settings. ✨ Finishing Touches🧪 Generate unit tests (beta)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
d62e9c7 to
75b65d2
Compare
|
@regression-perf |
Performance Regression Test ReportTriggered by: @Yongqi-Zhuo Results
Artifacts
|
ddf357e to
f3088ec
Compare
|
@regression-perf |
Performance Regression Test ReportTriggered by: @Yongqi-Zhuo Results
Artifacts
|
f3088ec to
704c267
Compare
|
@regression-perf |
Performance Regression Test ReportTriggered by: @Yongqi-Zhuo Results
Artifacts
|
588c33f to
c7585e6
Compare
|
@regression-perf |
Performance Regression Test ReportTriggered by: @Yongqi-Zhuo Results
Artifacts
|
|
@regression-perf |
Performance Regression Test ReportTriggered by: @Yongqi-Zhuo Results
Artifacts
|
c7585e6 to
4394ac9
Compare
There was a problem hiding this comment.
💡 Codex Review
Here are some automated review suggestions for this pull request.
Reviewed commit: 4394ac97bf
ℹ️ About Codex in GitHub
Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you
- Open a pull request for review
- Mark a draft as ready
- Comment "@codex review".
If Codex has suggestions, it will comment; otherwise it will react with 👍.
Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".
| if (cute::AsConst(tile_gbasis->shape[i]) > max_box_dim) { | ||
| DLOG(WARNING) << "The " << i | ||
| << "-th box dim > 256, fallback to normal copy"; | ||
| return fallback_to_normal("box dim > 256"); |
There was a problem hiding this comment.
Split oversized inner TMA boxes instead of falling back
When an unswizzled 2-D/ND T.tma_copy has a contiguous innermost tile larger than 256 elements, such as a 64x512 fp16 shared tile, tile_gbasis->shape[0] hits this branch and fallback_to_normal becomes a fatal error for explicit T.tma_copy. The previous lowering chunked the inner dimension to 256 and emitted a loop, which is valid because CUDA's limit is per boxDim, so this regresses valid TMA copies rather than just rejecting unsupported layouts.
Useful? React with 👍 / 👎.
| // tile. | ||
| int64_t inv_size = cute::AsConst(cute::Size(smem_plain_to_tile)); | ||
| int64_t tile_size = cute::AsConst(cute::Size(tile_to_smem_plain)); | ||
| ICHECK_EQ(inv_size, tile_size) |
There was a problem hiding this comment.
Fall back for gapped SMEM layouts instead of ICHECKing
For padded or gapped annotated shared layouts, for example a row-major tile with a padded row stride, RightInverse only covers the contiguous prefix, so this equality fails and aborts compilation. Earlier lowering explicitly fell back to normal copy for padded/unsupported layouts, and ordinary T.copy(..., prefer_instruction="tma") can still be lowered correctly without TMA, so this should use the existing fallback path rather than crashing the compile.
Useful? React with 👍 / 👎.
There was a problem hiding this comment.
Actionable comments posted: 2
🧹 Nitpick comments (2)
src/layout/cute_layout.cc (1)
34-38: 💤 Low valueConsider portability for
__builtin_ctzll.This is a GCC/Clang-specific builtin. If MSVC support is ever needed, a fallback would be required (e.g.,
_BitScanForward64). Given this is CUDA-focused code where GCC/Clang is standard, this is likely acceptable.🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@src/layout/cute_layout.cc` around lines 34 - 38, The Log2Exact function uses __builtin_ctzll, which is a GCC/Clang-specific builtin that is not supported by MSVC. To improve portability, add conditional compilation directives around the Log2Exact function implementation to use __builtin_ctzll for GCC/Clang and provide an alternative implementation using _BitScanForward64 for MSVC, or alternatively add a comment documenting that this function requires GCC/Clang and explaining what would be needed for MSVC support if that becomes necessary in the future.src/cuda/op/copy.cc (1)
1391-1399: Remove unusedmakeColumnMajorStridesfunction.The function is defined at lines 1391-1399 but never called anywhere in the codebase. This is dead code introduced in the PR.
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the rest with a brief reason, keep changes minimal, and validate. In `@src/cuda/op/copy.cc` around lines 1391 - 1399, The function `makeColumnMajorStrides` is defined but never called anywhere in the codebase, making it dead code. Remove the entire function definition including the function signature and its body. Search the codebase to confirm there are no calls to this function before deletion.
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
Inline comments:
In `@testing/python/layout/test_tilelang_cute.py`:
- Line 48: The function `forward` uses `vars` as a parameter name, which shadows
Python's built-in `vars()` function and is flagged by Ruff linting. Rename the
`vars` parameter in the `forward` function signature to a non-conflicting name
like `args` or `arguments`. Additionally, there is an ambiguous identifier `I`
(referenced at line 117) that should also be renamed to a more descriptive
identifier to avoid Ruff linting errors. These changes will resolve the CI
linting blockers.
In `@tilelang/cuda/intrinsics/layout/mma_layout.py`:
- Around line 234-240: The code at line 234 calculates swizzle_vectors by
dividing swizzle_bytes by 16, but it does not validate that swizzle_bytes
contains an allowed value before this division. Add validation logic before the
swizzle_vectors calculation to ensure swizzle_bytes is one of the supported
values (such as 32, 64, or 128 bytes) that map to valid swizzle vector counts of
2, 4, or 8. If an unsupported swizzle_bytes value is encountered, raise an
appropriate exception with a descriptive error message. This prevents division
by zero or invalid swizzle mappings in the subsequent calculations.
---
Nitpick comments:
In `@src/cuda/op/copy.cc`:
- Around line 1391-1399: The function `makeColumnMajorStrides` is defined but
never called anywhere in the codebase, making it dead code. Remove the entire
function definition including the function signature and its body. Search the
codebase to confirm there are no calls to this function before deletion.
In `@src/layout/cute_layout.cc`:
- Around line 34-38: The Log2Exact function uses __builtin_ctzll, which is a
GCC/Clang-specific builtin that is not supported by MSVC. To improve
portability, add conditional compilation directives around the Log2Exact
function implementation to use __builtin_ctzll for GCC/Clang and provide an
alternative implementation using _BitScanForward64 for MSVC, or alternatively
add a comment documenting that this function requires GCC/Clang and explaining
what would be needed for MSVC support if that becomes necessary in the future.
🪄 Autofix (Beta)
Fix all unresolved CodeRabbit comments on this PR:
- Push a commit to this branch (recommended)
- Create a new PR with the fixes
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro
Run ID: 3be2ebae-8119-48e9-841e-b9e855d81460
📒 Files selected for processing (11)
examples/gemm/example_gemm_intrinsics.pysrc/cuda/op/copy.ccsrc/cuda/transform/producer_consumer_ws.ccsrc/layout/cute_layout.ccsrc/layout/cute_layout.htesting/python/kernel/test_tilelang_kernel_gemm_simt.pytesting/python/layout/test_tilelang_cute.pytilelang/cuda/intrinsics/layout/mma_layout.pytilelang/layout/__init__.pytilelang/layout/_cute_ffi_api.pytilelang/layout/cute.py
| apply(x) = x ^ ((x & yyy) >> s_shift), with yyy = mask << (m_base + s_shift) | ||
| and mask = (1 << b_bits) - 1.""" | ||
|
|
||
| def forward(*vars): |
There was a problem hiding this comment.
Rename lint-blocking identifiers (vars, I).
forward(*vars) shadows a Python builtin, and I is ambiguous. Ruff flags both as errors, so this can block linted CI.
Suggested rename patch
- def forward(*vars):
- addr = intermediate_fn(*vars)
+ def forward(*coords):
+ addr = intermediate_fn(*coords)
@@
- I = cute.make_identity_layout((2, (2, 2)))
- assert I.shape == (2, (2, 2))
- assert [I(c) for c in range(8)] == [(c % 2, (c // 2 % 2, c // 4)) for c in range(8)]
+ identity_layout = cute.make_identity_layout((2, (2, 2)))
+ assert identity_layout.shape == (2, (2, 2))
+ assert [identity_layout(c) for c in range(8)] == [(c % 2, (c // 2 % 2, c // 4)) for c in range(8)]Also applies to: 117-117
🧰 Tools
🪛 Ruff (0.15.17)
[error] 48-48: Function argument vars is shadowing a Python builtin
(A002)
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@testing/python/layout/test_tilelang_cute.py` at line 48, The function
`forward` uses `vars` as a parameter name, which shadows Python's built-in
`vars()` function and is flagged by Ruff linting. Rename the `vars` parameter in
the `forward` function signature to a non-conflicting name like `args` or
`arguments`. Additionally, there is an ambiguous identifier `I` (referenced at
line 117) that should also be renamed to a more descriptive identifier to avoid
Ruff linting errors. These changes will resolve the CI linting blockers.
Source: Linters/SAST tools
| swizzle_vectors = int(swizzle_bytes) // 16 | ||
| col_idx_16B = col_idx // elem_per_16B | ||
| col_idx_in_16B = col_idx % elem_per_16B | ||
| new_col_idx_16B = col_idx_16B ^ (row_idx % (swizzle_bytes // 16)) | ||
| return row_idx, ana.simplify(new_col_idx_16B * elem_per_16B + col_idx_in_16B) | ||
| col_tile = col_idx_16B // swizzle_vectors | ||
| c = col_idx_16B % swizzle_vectors | ||
| src = (row_idx % 8) // (8 // swizzle_vectors) | ||
| swizzled_col = (c ^ src) * elem_per_16B + col_idx_in_16B |
There was a problem hiding this comment.
Validate swizzle_bytes domain before deriving swizzle_vectors.
Line 234/239 assume swizzle_bytes maps to valid vector counts (2/4/8), but unsupported values can crash (// 0) or produce invalid swizzle mappings. Add an explicit guard for allowed swizzle sizes and bounds.
Suggested patch
def get_swizzle_layout(row_idx, col_idx, row_size, dtype: DataType | str, swizzle_bytes=None):
@@
if swizzle_bytes is None:
swizzle_bytes = min(128, row_bytes)
+ if swizzle_bytes not in (32, 64, 128):
+ raise ValueError("swizzle_bytes must be one of {32, 64, 128}.")
+ if swizzle_bytes > row_bytes:
+ raise ValueError("swizzle_bytes cannot exceed row size in bytes.")
@@
- swizzle_vectors = int(swizzle_bytes) // 16
+ swizzle_vectors = swizzle_bytes // 16🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@tilelang/cuda/intrinsics/layout/mma_layout.py` around lines 234 - 240, The
code at line 234 calculates swizzle_vectors by dividing swizzle_bytes by 16, but
it does not validate that swizzle_bytes contains an allowed value before this
division. Add validation logic before the swizzle_vectors calculation to ensure
swizzle_bytes is one of the supported values (such as 32, 64, or 128 bytes) that
map to valid swizzle vector counts of 2, 4, or 8. If an unsupported
swizzle_bytes value is encountered, raise an appropriate exception with a
descriptive error message. This prevents division by zero or invalid swizzle
mappings in the subsequent calculations.
|
@regression-perf |
4394ac9 to
3cc4648
Compare
There was a problem hiding this comment.
Actionable comments posted: 2
🤖 Prompt for all review comments with AI agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
Inline comments:
In `@src/layout/cute_layout.cc`:
- Around line 1458-1490: The code currently commits to the first contiguous run
found in s_at and returns early with std::nullopt if RecoverPlainLayout fails
for that candidate, but it should instead enumerate all candidate contiguous
runs and try each one. Refactor the swizzle candidate selection logic to collect
all valid contiguous runs of target bits that share an s_shift value, then
iterate through each candidate swizzle built from those runs (including
Swizzle::Identity as a fallback), testing each with RecoverPlainLayout and
returning the first one that succeeds. Only return std::nullopt if no candidate
passes validation, avoiding false rejections of valid affine layouts with
non-power-of-two strides.
- Around line 973-1029: The AddrProbe class performs address arithmetic with
narrowing from int64_t to int32_t without overflow checks, risking undefined
behavior on large constant layouts. In the AddrProbe constructor, accumulate
out_strides_ using int64_t (keep acc as int64_t throughout the loop), then add
explicit range checks to verify each accumulated stride value fits in the
int32_t range before storing it in out_strides_. Similarly, in the operator()
method, before performing the address calculation with addr +=
static_cast<int32_t>(*val) * out_strides_[d], add range checks to ensure both
val and the resulting multiplication fit within int32_t bounds, and verify the
accumulated addr doesn't overflow int32_t range before returning it.
🪄 Autofix (Beta)
Fix all unresolved CodeRabbit comments on this PR:
- Push a commit to this branch (recommended)
- Create a new PR with the fixes
ℹ️ Review info
⚙️ Run configuration
Configuration used: defaults
Review profile: CHILL
Plan: Pro
Run ID: 765628c6-de2e-4b4e-8a7f-82e095bba332
📒 Files selected for processing (11)
examples/gemm/example_gemm_intrinsics.pysrc/cuda/op/copy.ccsrc/cuda/transform/producer_consumer_ws.ccsrc/layout/cute_layout.ccsrc/layout/cute_layout.htesting/python/kernel/test_tilelang_kernel_gemm_simt.pytesting/python/layout/test_tilelang_cute.pytilelang/cuda/intrinsics/layout/mma_layout.pytilelang/layout/__init__.pytilelang/layout/_cute_ffi_api.pytilelang/layout/cute.py
✅ Files skipped from review due to trivial changes (1)
- examples/gemm/example_gemm_intrinsics.py
🚧 Files skipped from review as they are similar to previous changes (8)
- tilelang/layout/init.py
- src/cuda/transform/producer_consumer_ws.cc
- tilelang/layout/_cute_ffi_api.py
- testing/python/kernel/test_tilelang_kernel_gemm_simt.py
- tilelang/cuda/intrinsics/layout/mma_layout.py
- src/cuda/op/copy.cc
- src/layout/cute_layout.h
- tilelang/layout/cute.py
| // When processing shared tensor layout, all the shapes and strides are int32_t. | ||
| // Also, we need to make sure not to mix with int64_t, because otherwise there | ||
| // would be a lot of conversion nodes that prevent the analyzer from cancelling | ||
| // out the terms. | ||
| PrimExpr MakeI32(int x) { return IntImm(DataType::Int(32), x); } | ||
|
|
||
| // Probes a TileLang layout's row-major linearized physical address A(x). | ||
| // The constant input shape, output strides, forward-index expressions, and | ||
| // input placeholders are parsed once in the constructor. operator() folds | ||
| // concrete integer coordinates to a constant address; Symbolic() substitutes | ||
| // fresh variables for the equivalence proof. valid() is false when any extent | ||
| // is non-constant. | ||
| class AddrProbe { | ||
| public: | ||
| explicit AddrProbe(const tvm::tl::Layout &layout) { | ||
| ICHECK(layout.defined()); | ||
| for (const auto &e : layout->InputShape()) { | ||
| auto c = as_const_int(e); | ||
| ICHECK(c) << "InputShape extent " << e | ||
| << " of a shared tensor layout must be constant"; | ||
| shape_.push_back(*c); | ||
| } | ||
| std::vector<int32_t> out_sizes; | ||
| for (const auto &e : layout->OutputShape()) { | ||
| auto c = as_const_int(e); | ||
| ICHECK(c) << "OutputShape extent " << e | ||
| << " of a shared tensor layout must be constant"; | ||
| out_sizes.push_back(*c); | ||
| } | ||
| out_strides_.assign(out_sizes.size(), 0); | ||
| int32_t acc = 1; | ||
| for (int64_t d = static_cast<int64_t>(out_sizes.size()) - 1; d >= 0; --d) { | ||
| out_strides_[d] = acc; | ||
| acc *= out_sizes[d]; | ||
| } | ||
| forward_index_ = layout->GetForwardIndex(); | ||
| ICHECK_EQ(forward_index_.size(), out_strides_.size()); | ||
| ICHECK_EQ(layout->InputDim(), shape_.size()); | ||
| for (size_t k = 0; k < shape_.size(); ++k) | ||
| placeholders_.push_back(InputPlaceholder(k)); | ||
| } | ||
|
|
||
| const std::vector<int32_t> &shape() const { return shape_; } | ||
|
|
||
| // Concrete address A(coords); nullopt if it does not fold to a constant. | ||
| std::optional<int32_t> operator()(const std::vector<int32_t> &coords) const { | ||
| Map<Var, PrimExpr> vmap; | ||
| for (size_t k = 0; k < coords.size(); ++k) | ||
| vmap.Set(placeholders_[k], MakeI32(coords[k])); | ||
| int32_t addr = 0; | ||
| for (size_t d = 0; d < forward_index_.size(); ++d) { | ||
| PrimExpr e = Substitute(forward_index_[d], vmap); | ||
| std::optional<int64_t> val = EvalConstExpr(e); | ||
| if (!val) | ||
| return std::nullopt; | ||
| addr += static_cast<int32_t>(*val) * out_strides_[d]; | ||
| } |
There was a problem hiding this comment.
Reject out-of-range address arithmetic before narrowing to int32_t.
as_const_int yields int64_t, but extents, output strides, addr, and MakeI32(offset) are narrowed or accumulated as signed int32_t. Large-but-constant layouts can overflow before recovery returns None, which risks UB or proving a wrapped address map. Use checked int64_t accumulation and explicit range checks before creating int32 expressions.
Guard the narrowing points
-PrimExpr MakeI32(int x) { return IntImm(DataType::Int(32), x); }
+int32_t CheckedI32(int64_t value, const char *what) {
+ ICHECK(value >= std::numeric_limits<int32_t>::min() &&
+ value <= std::numeric_limits<int32_t>::max())
+ << what << " is outside int32 range: " << value;
+ return static_cast<int32_t>(value);
+}
+
+PrimExpr MakeI32(int64_t x) {
+ return IntImm(DataType::Int(32), CheckedI32(x, "int32 PrimExpr constant"));
+}
@@
- shape_.push_back(*c);
+ shape_.push_back(CheckedI32(*c, "InputShape extent"));
@@
- out_sizes.push_back(*c);
+ out_sizes.push_back(CheckedI32(*c, "OutputShape extent"));
@@
- int32_t acc = 1;
+ int64_t acc = 1;
for (int64_t d = static_cast<int64_t>(out_sizes.size()) - 1; d >= 0; --d) {
- out_strides_[d] = acc;
+ out_strides_[d] = CheckedI32(acc, "OutputShape stride");
+ ICHECK(out_sizes[d] > 0) << "OutputShape extent must be positive";
+ ICHECK_LE(acc, std::numeric_limits<int32_t>::max() / out_sizes[d])
+ << "OutputShape product exceeds int32 range";
acc *= out_sizes[d];
}
@@
- int32_t addr = 0;
+ int64_t addr = 0;
@@
- addr += static_cast<int32_t>(*val) * out_strides_[d];
+ addr += *val * static_cast<int64_t>(out_strides_[d]);
@@
- return addr;
+ return CheckedI32(addr, "linearized address");Also applies to: 1314-1320
🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@src/layout/cute_layout.cc` around lines 973 - 1029, The AddrProbe class
performs address arithmetic with narrowing from int64_t to int32_t without
overflow checks, risking undefined behavior on large constant layouts. In the
AddrProbe constructor, accumulate out_strides_ using int64_t (keep acc as
int64_t throughout the loop), then add explicit range checks to verify each
accumulated stride value fits in the int32_t range before storing it in
out_strides_. Similarly, in the operator() method, before performing the address
calculation with addr += static_cast<int32_t>(*val) * out_strides_[d], add range
checks to ensure both val and the resulting multiplication fit within int32_t
bounds, and verify the accumulated addr doesn't overflow int32_t range before
returning it.
| std::map<int, int> s_at; // swizzle target bit -> s_shift. | ||
| for (uint64_t col : cols) { | ||
| if (__builtin_popcountll(col) != 2) | ||
| continue; | ||
| int lo = __builtin_ctzll(col); | ||
| int hi = 63 - __builtin_clzll(col); | ||
| if (!(weight1 & (uint64_t(1) << lo))) | ||
| continue; // two-bit column from a plain stride, not a swizzle. | ||
| if (s_at.count(lo)) | ||
| return std::nullopt; // two sources hitting one target: not a Swizzle. | ||
| s_at.emplace(lo, hi - lo); | ||
| } | ||
| Swizzle swizzle = Swizzle::Identity(); | ||
| if (!s_at.empty()) { | ||
| // Take the lowest contiguous run of target bits that share one s_shift. | ||
| int m_base = s_at.begin()->first; | ||
| int s_shift = s_at.begin()->second; | ||
| int b_bits = 0; | ||
| for (auto it = s_at.find(m_base + b_bits); | ||
| it != s_at.end() && it->second == s_shift; | ||
| it = s_at.find(m_base + b_bits)) | ||
| ++b_bits; | ||
| if (s_shift < b_bits) | ||
| return std::nullopt; // source and target bit regions must not overlap. | ||
| swizzle = Swizzle(b_bits, m_base, s_shift); | ||
| } | ||
|
|
||
| // The base offset is Sw(A(0)), since Sw is an involution; recover the plain | ||
| // layout under that swizzle and offset. | ||
| int64_t offset = swizzle->Apply(*A0); | ||
| Optional<Layout> plain = RecoverPlainLayout(A, swizzle, offset); | ||
| if (!plain.defined()) | ||
| return std::nullopt; |
There was a problem hiding this comment.
Try proven swizzle candidates instead of committing to the lowest two-bit column.
The detector can false-reject valid affine layouts: an ordinary stride like 3 produces a two-bit column {0,1}, and if another mode has stride 1, weight1 makes it look like Sw<1,0,1>. The selected swizzle then fails RecoverPlainLayout, but the code returns None instead of falling back to identity or trying later real swizzle runs. Avoid returning before proof on duplicate/ambiguous candidates, enumerate candidate contiguous runs, and return the first one that RecoverPlainLayout proves; try identity as the fallback.
Suggested regression shape
def test_plain_nonpow2_stride_does_not_force_swizzle():
mode = cute.ComposedLayout.from_tilelang(Layout((2, 2), lambda i, j: i * 3 + j))
assert mode is not None
assert not mode.swizzle.is_swizzled
_assert_struct(mode.layout, (2, 2), (3, 1))🤖 Prompt for AI Agents
Verify each finding against current code. Fix only still-valid issues, skip the
rest with a brief reason, keep changes minimal, and validate.
In `@src/layout/cute_layout.cc` around lines 1458 - 1490, The code currently
commits to the first contiguous run found in s_at and returns early with
std::nullopt if RecoverPlainLayout fails for that candidate, but it should
instead enumerate all candidate contiguous runs and try each one. Refactor the
swizzle candidate selection logic to collect all valid contiguous runs of target
bits that share an s_shift value, then iterate through each candidate swizzle
built from those runs (including Swizzle::Identity as a fallback), testing each
with RecoverPlainLayout and returning the first one that succeeds. Only return
std::nullopt if no candidate passes validation, avoiding false rejections of
valid affine layouts with non-power-of-two strides.
Performance Regression Test ReportTriggered by: @Yongqi-Zhuo Results
Artifacts
|
3cc4648 to
847eeda
Compare
|
@regression-perf |
Performance Regression Test ReportTriggered by: @Yongqi-Zhuo Results
Artifacts
|
Summary by CodeRabbit
Release Notes