[Feature] Support cluster launch, query, synchronization and barrier operations - #1874
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! 🚀 |
|
Note Reviews pausedIt looks like this branch is under active development. To avoid overwhelming you with review comments due to an influx of new commits, CodeRabbit has automatically paused this review. You can configure this behavior by changing the Use the following commands to manage reviews:
Use the checkboxes below for quick actions:
📝 WalkthroughWalkthroughAdds end-to-end SM90+ CUDA cluster-dimensions support: new Kernel API arg, propagation of Changes
Sequence DiagramsequenceDiagram
participant User as User Code
participant KernelAPI as Kernel API
participant Transform as Transform Passes
participant HostCodegen as Host Codegen / JIT Wrapper
participant CUDARuntime as CUDA Runtime
User->>KernelAPI: Call Kernel(cluster_dims=[x,y,z])
KernelAPI->>KernelAPI: Normalize and attach cluster_dims attr
KernelAPI->>Transform: Emit PrimFunc with cluster_dims
rect rgba(100,150,200,0.5)
Note over Transform: Propagate cluster_dims through transforms
Transform->>Transform: split_host_device: read & SetClusterDims
Transform->>Transform: lower_opaque_block: capture annotation
Transform->>Transform: lower_device_kernel_launch: prepend launch args/params
end
Transform->>HostCodegen: Device function with cluster_dims attr
HostCodegen->>HostCodegen: Choose launch path
alt cluster_dims present
HostCodegen->>CUDARuntime: Emit cluster launch (KERNEL_CLUSTER_LAUNCH)
else
HostCodegen->>CUDARuntime: Emit standard launch (KERNEL_LAUNCH)
end
CUDARuntime->>CUDARuntime: Extract launch tags (skip duplicate cluster tags)
CUDARuntime->>CUDARuntime: Configure SM90+ cluster launch
Estimated code review effort🎯 4 (Complex) | ⏱️ ~45 minutes Possibly related PRs
Suggested labels
Suggested reviewers
Poem
🚥 Pre-merge checks | ✅ 2 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (2 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 |
|
Also plan to add assertation for cuda sm90+ and backend execution, but not sure where is the best position to insert this. |
There was a problem hiding this comment.
Actionable comments posted: 2
Caution
Some comments are outside the diff and can’t be posted inline due to platform limitations.
⚠️ Outside diff range comments (2)
src/transform/lower_device_kernel_launch.cc (1)
85-124:⚠️ Potential issue | 🟠 MajorAdd length validation before indexing cluster_dims.
dims[0..2]is read unconditionally; a malformed attr can crash or miscompile. Add a size check (and optionally a >0 check) before indexing.✅ Suggested fix
if (collector.info_.cluster_dims.defined()) { auto dims = collector.info_.cluster_dims.value(); + ICHECK_EQ(dims.size(), 3) << "cluster_dims must have 3 elements"; Array<PrimExpr> new_launch_args = { PrimExpr(dims[0]), PrimExpr(dims[1]), PrimExpr(dims[2])};🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@src/transform/lower_device_kernel_launch.cc` around lines 85 - 124, The code reads collector.info_.cluster_dims.value() into dims and indexes dims[0..2] without checking length; add a guard that ensures collector.info_.cluster_dims.defined() && dims.size() >= 3 (and optionally each dim > 0) before constructing new_launch_args/new_launch_params; if the check fails, log an error via LOG(ERROR) or skip prepending cluster dims (or return/fail gracefully) so malformed cluster_dims from func->GetAttr<Array<Integer>> cannot cause out-of-bounds access in the block that builds new_launch_args/new_launch_params.tilelang/jit/adapter/wrapper.py (1)
475-510:⚠️ Potential issue | 🟡 MinorValidate cluster_dims length before use in format() calls.
KERNEL_CLUSTER_LAUNCH_FUNC_CODErequires exactly 3 values. If attrs are malformed,.format(*cluster_dims)can raise or mislaunch. Consider padding (len<3) and rejecting len>3 at parse time.✅ Suggested fix
if "cluster_dims" in attrs: # Extract cluster dimensions for SM90+ cluster launch cluster_dims_attr = attrs["cluster_dims"] - cluster_dims = [int(cluster_dims_attr[i]) for i in range(len(cluster_dims_attr))] + cluster_dims = [int(cluster_dims_attr[i]) for i in range(len(cluster_dims_attr))] + if len(cluster_dims) > 3: + raise ValueError("cluster_dims must have at most 3 elements") + cluster_dims = cluster_dims + [1] * (3 - len(cluster_dims))Also applies to: 515-518, 615-616
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@tilelang/jit/adapter/wrapper.py` around lines 475 - 510, cluster_dims extracted from attrs may have length != 3 which will break or misconfigure the KERNEL_CLUSTER_LAUNCH_FUNC_CODE.format(*cluster_dims) calls; validate cluster_dims when building cluster_dims_map and before every .format usage (e.g., where KERNEL_CLUSTER_LAUNCH_FUNC_CODE is used around lines noted) by: 1) rejecting and logging/raising if len(cluster_dims) > 3, 2) padding with 1s to length 3 when len < 3, and 3) ensure any downstream code that reads cluster_dims (cluster_dims_map and the format sites) receives a guaranteed 3-element list so .format(*cluster_dims) is safe.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@src/transform/lower_opaque_block.cc`:
- Around line 58-60: The code currently applies cluster_dims_ to the attribute
bag (using WithAttr and variable f) without validating multiple or malformed
annotations; update the logic that sets cluster_dims_ (references:
lower.cluster_dims_, WithAttr, and the code paths that attach the "cluster_dims"
attribute) to (1) detect and error/log if more than one annotation was provided
(or otherwise enforce a clear precedence rule instead of silently overwriting),
and (2) validate the length/shape of the cluster_dims_ value before attaching
(reject or normalize values that don't match expected rank/length). Add
consistent checks where cluster_dims_ is applied (the other analogous places
that attach "cluster_dims") so all code paths perform the same validation and
fail fast on conflicts or malformed inputs.
In `@tilelang/language/kernel.py`:
- Around line 318-327: The cluster_dims handling currently accepts lists/tuples
longer than 3 silently; update the block that processes cluster_dims (the branch
referencing cluster_dims, isinstance checks, and attrs["cluster_dims"]) to
validate that if cluster_dims is a list or tuple its length is <= 3 and if it's
longer raise a ValueError; keep the existing padding behavior (extend with [1]
to length 3) only for lengths 1–3, handle an int by converting to [int,1,1], and
tighten the ValueError message to clearly state the accepted types and maximum
length (e.g. "cluster_dims must be an int or a list/tuple of up to 3 integers").
---
Outside diff comments:
In `@src/transform/lower_device_kernel_launch.cc`:
- Around line 85-124: The code reads collector.info_.cluster_dims.value() into
dims and indexes dims[0..2] without checking length; add a guard that ensures
collector.info_.cluster_dims.defined() && dims.size() >= 3 (and optionally each
dim > 0) before constructing new_launch_args/new_launch_params; if the check
fails, log an error via LOG(ERROR) or skip prepending cluster dims (or
return/fail gracefully) so malformed cluster_dims from
func->GetAttr<Array<Integer>> cannot cause out-of-bounds access in the block
that builds new_launch_args/new_launch_params.
In `@tilelang/jit/adapter/wrapper.py`:
- Around line 475-510: cluster_dims extracted from attrs may have length != 3
which will break or misconfigure the
KERNEL_CLUSTER_LAUNCH_FUNC_CODE.format(*cluster_dims) calls; validate
cluster_dims when building cluster_dims_map and before every .format usage
(e.g., where KERNEL_CLUSTER_LAUNCH_FUNC_CODE is used around lines noted) by: 1)
rejecting and logging/raising if len(cluster_dims) > 3, 2) padding with 1s to
length 3 when len < 3, and 3) ensure any downstream code that reads cluster_dims
(cluster_dims_map and the format sites) receives a guaranteed 3-element list so
.format(*cluster_dims) is safe.
ℹ️ Review info
Configuration used: defaults
Review profile: CHILL
Plan: Pro
📒 Files selected for processing (7)
examples/gemm_sm100/gemm_tcgen5mma_ws_persistent.pysrc/target/rt_mod_cuda.ccsrc/transform/lower_device_kernel_launch.ccsrc/transform/lower_opaque_block.ccsrc/transform/split_host_device.cctilelang/jit/adapter/wrapper.pytilelang/language/kernel.py
| if (lower.cluster_dims_.has_value()) { | ||
| f = WithAttr(std::move(f), "cluster_dims", lower.cluster_dims_.value()); | ||
| } |
There was a problem hiding this comment.
Guard against conflicting or malformed cluster_dims annotations.
If multiple annotations appear, the last one wins silently. Adding a consistency/length check will prevent accidental overrides and invalid launches.
✅ Suggested fix
} else if (key == "cluster_dims") {
if (auto arr = kv.second.try_cast<Array<Integer>>()) {
- cluster_dims_ = arr.value();
+ auto value = arr.value();
+ ICHECK_EQ(value.size(), 3) << "cluster_dims must have 3 elements";
+ if (cluster_dims_.defined()) {
+ auto prev = cluster_dims_.value();
+ ICHECK_EQ(prev.size(), value.size()) << "Conflicting cluster_dims annotations";
+ for (int i = 0; i < static_cast<int>(value.size()); ++i) {
+ ICHECK_EQ(prev[i], value[i]) << "Conflicting cluster_dims annotations";
+ }
+ }
+ cluster_dims_ = value;
} else {
LOG(FATAL) << "Expected `" << "cluster_dims"
<< "` to be an Array<Integer>, but got "
<< kv.second.GetTypeKey();
}Also applies to: 260-267, 305-306
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@src/transform/lower_opaque_block.cc` around lines 58 - 60, The code currently
applies cluster_dims_ to the attribute bag (using WithAttr and variable f)
without validating multiple or malformed annotations; update the logic that sets
cluster_dims_ (references: lower.cluster_dims_, WithAttr, and the code paths
that attach the "cluster_dims" attribute) to (1) detect and error/log if more
than one annotation was provided (or otherwise enforce a clear precedence rule
instead of silently overwriting), and (2) validate the length/shape of the
cluster_dims_ value before attaching (reject or normalize values that don't
match expected rank/length). Add consistent checks where cluster_dims_ is
applied (the other analogous places that attach "cluster_dims") so all code
paths perform the same validation and fail fast on conflicts or malformed
inputs.
| if cluster_dims is not None: | ||
| if isinstance(cluster_dims, (list, tuple)): | ||
| cluster_dims = list(cluster_dims) + [1] * (3 - len(cluster_dims)) | ||
| elif isinstance(cluster_dims, int): | ||
| cluster_dims = [cluster_dims, 1, 1] | ||
| else: | ||
| raise ValueError("cluster_dims must be a list or tuple of integers") | ||
|
|
||
| if cluster_dims != [1, 1, 1]: | ||
| attrs["cluster_dims"] = cluster_dims |
There was a problem hiding this comment.
Validate cluster_dims length (≤3) and tighten error messaging.
Lists/tuples longer than 3 are silently preserved, which can lead to downstream mismatch or indexing failures. Consider rejecting >3 and clarifying the accepted types.
✅ Suggested fix
- if cluster_dims is not None:
- if isinstance(cluster_dims, (list, tuple)):
- cluster_dims = list(cluster_dims) + [1] * (3 - len(cluster_dims))
- elif isinstance(cluster_dims, int):
- cluster_dims = [cluster_dims, 1, 1]
- else:
- raise ValueError("cluster_dims must be a list or tuple of integers")
+ if cluster_dims is not None:
+ if isinstance(cluster_dims, (list, tuple)):
+ cluster_dims = list(cluster_dims)
+ if len(cluster_dims) > 3:
+ raise ValueError("cluster_dims must have at most 3 elements")
+ cluster_dims = cluster_dims + [1] * (3 - len(cluster_dims))
+ elif isinstance(cluster_dims, int):
+ cluster_dims = [cluster_dims, 1, 1]
+ else:
+ raise ValueError("cluster_dims must be an int or a list/tuple of ints")🧰 Tools
🪛 Ruff (0.15.2)
[warning] 324-324: Avoid specifying long messages outside the exception class
(TRY003)
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@tilelang/language/kernel.py` around lines 318 - 327, The cluster_dims
handling currently accepts lists/tuples longer than 3 silently; update the block
that processes cluster_dims (the branch referencing cluster_dims, isinstance
checks, and attrs["cluster_dims"]) to validate that if cluster_dims is a list or
tuple its length is <= 3 and if it's longer raise a ValueError; keep the
existing padding behavior (extend with [1] to length 3) only for lengths 1–3,
handle an int by converting to [int,1,1], and tighten the ValueError message to
clearly state the accepted types and maximum length (e.g. "cluster_dims must be
an int or a list/tuple of up to 3 integers").
There was a problem hiding this comment.
Actionable comments posted: 1
♻️ Duplicate comments (1)
src/transform/lower_opaque_block.cc (1)
260-267:⚠️ Potential issue | 🟡 MinorAdd a size check to ensure
cluster_dimshas exactly 3 elements.Downstream code (e.g.,
lower_device_kernel_launch.ccline 109) indexesdims[0],dims[1],dims[2]unconditionally. If a malformed annotation provides fewer than 3 elements, this will cause an out-of-bounds access.Proposed fix
} else if (key == "cluster_dims") { if (auto arr = kv.second.try_cast<Array<Integer>>()) { - cluster_dims_ = arr.value(); + auto value = arr.value(); + ICHECK_EQ(value.size(), 3) + << "cluster_dims annotation must have exactly 3 elements, got " + << value.size(); + cluster_dims_ = value; } else {🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@src/transform/lower_opaque_block.cc` around lines 260 - 267, The code accepts a "cluster_dims" Array<Integer> into cluster_dims_ but lacks a size check; add a guard after assigning cluster_dims_ (in the block handling key == "cluster_dims" in lower_opaque_block.cc) that verifies cluster_dims_.size() == 3 and emits a fatal/log error if not, so downstream code (e.g., lower_device_kernel_launch.cc which indexes dims[0..2]) cannot experience out-of-bounds access; keep the existing type check and only perform this size validation immediately after setting cluster_dims_.
🧹 Nitpick comments (2)
tilelang/jit/adapter/wrapper.py (2)
176-197: MovecudaFuncSetAttribute(…NonPortableClusterSizeAllowed…)to the one-time init function.
cudaFuncSetAttributeis called on every kernel invocation (line 190) but the attribute is a static property of the kernel. Moving it toget_init_func()(alongside the existing dynamic shared memory attribute setup) avoids a redundant driver call on each launch.Suggested approach
In
get_init_func, add a cluster attribute block similar to the dynamic shared memory one:def get_init_func(self): call_str = """""" for function_name, dynamic_smem_buf in self.dynamic_smem_buf.items(): if dynamic_smem_buf is not None: call_str += PREDEF_ATTRIBUTE_SET_DYNAMIC_MEMORY.format(function_name, dynamic_smem_buf) + for function_name, cluster_dims in self.cluster_dims.items(): + if cluster_dims is not None: + call_str += ( + f' cudaError_t cluster_result_{function_name} = ' + f'cudaFuncSetAttribute({function_name}, ' + f'cudaFuncAttributeNonPortableClusterSizeAllowed, 1);\n' + f' if (cluster_result_{function_name} != cudaSuccess) {{\n' + f' snprintf(error_buf, ERROR_BUF_SIZE, ' + f'"Failed to set cluster attribute for {function_name}: %s", ' + f'cudaGetErrorString(cluster_result_{function_name}));\n' + f' return -1;\n' + f' }}\n' + ) init_funcs = PREDEF_INIT_FUNC.format(call_str) return init_funcsThen remove the
cudaFuncSetAttribute+ error check fromKERNEL_CLUSTER_LAUNCH_FUNC_CODE.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@tilelang/jit/adapter/wrapper.py` around lines 176 - 197, The kernel cluster attribute is being set on every launch in KERNEL_CLUSTER_LAUNCH_FUNC_CODE; move that one-time call into the init routine: add a cluster-attribute setup in get_init_func() (parallel to the existing dynamic shared memory attribute setup) that calls cudaFuncSetAttribute(func, cudaFuncAttributeNonPortableClusterSizeAllowed, 1) with proper error handling and logging, and then remove the cudaFuncSetAttribute + error check block from KERNEL_CLUSTER_LAUNCH_FUNC_CODE so kernel launches only configure the launch config and do not re-set the static attribute.
610-610:cluster_dimsinfunction_informationsis stored but never consumed.
create_dispatch_funcreadsblock_info,grid_info,dynamic_smem_buf, andfunction_paramsfromfunction_info, but accessescluster_dimsdirectly viaself.cluster_dims[function_name](line 363). The entry on line 610 is dead data. Either read it fromfunction_infoconsistently (aligning with the other fields) or remove it to avoid confusion.🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@tilelang/jit/adapter/wrapper.py` at line 610, function_informations stores "cluster_dims" but create_dispatch_func currently ignores it and uses self.cluster_dims[function_name]; update create_dispatch_func to read cluster_dims from the function_info dict (the same way it reads "block_info", "grid_info", "dynamic_smem_buf", "function_params") i.e. use function_info.get("cluster_dims", None) instead of self.cluster_dims[function_name]; this makes the stored entry at "cluster_dims" meaningful and removes the inconsistent direct access to self.cluster_dims.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@tilelang/jit/adapter/wrapper.py`:
- Around line 492-495: The parsed cluster_dims_attr is not checked for length
and may be shorter than 3 causing downstream IndexError when used (see
cluster_dims, cluster_dims_attr and self.cluster_dims[function_name]); add a
validation right after building cluster_dims to ensure len(cluster_dims) == 3
and raise a clear ValueError or AssertionError with a message referencing the
offending attribute and function_name (e.g., "expected 3 cluster_dims for
function {function_name}, got {len(cluster_dims)}") so the error is reported at
the origin before any .format or unpacking occurs.
---
Duplicate comments:
In `@src/transform/lower_opaque_block.cc`:
- Around line 260-267: The code accepts a "cluster_dims" Array<Integer> into
cluster_dims_ but lacks a size check; add a guard after assigning cluster_dims_
(in the block handling key == "cluster_dims" in lower_opaque_block.cc) that
verifies cluster_dims_.size() == 3 and emits a fatal/log error if not, so
downstream code (e.g., lower_device_kernel_launch.cc which indexes dims[0..2])
cannot experience out-of-bounds access; keep the existing type check and only
perform this size validation immediately after setting cluster_dims_.
---
Nitpick comments:
In `@tilelang/jit/adapter/wrapper.py`:
- Around line 176-197: The kernel cluster attribute is being set on every launch
in KERNEL_CLUSTER_LAUNCH_FUNC_CODE; move that one-time call into the init
routine: add a cluster-attribute setup in get_init_func() (parallel to the
existing dynamic shared memory attribute setup) that calls
cudaFuncSetAttribute(func, cudaFuncAttributeNonPortableClusterSizeAllowed, 1)
with proper error handling and logging, and then remove the cudaFuncSetAttribute
+ error check block from KERNEL_CLUSTER_LAUNCH_FUNC_CODE so kernel launches only
configure the launch config and do not re-set the static attribute.
- Line 610: function_informations stores "cluster_dims" but create_dispatch_func
currently ignores it and uses self.cluster_dims[function_name]; update
create_dispatch_func to read cluster_dims from the function_info dict (the same
way it reads "block_info", "grid_info", "dynamic_smem_buf", "function_params")
i.e. use function_info.get("cluster_dims", None) instead of
self.cluster_dims[function_name]; this makes the stored entry at "cluster_dims"
meaningful and removes the inconsistent direct access to self.cluster_dims.
ℹ️ Review info
Configuration used: defaults
Review profile: CHILL
Plan: Pro
📒 Files selected for processing (3)
src/transform/lower_device_kernel_launch.ccsrc/transform/lower_opaque_block.cctilelang/jit/adapter/wrapper.py
| if "cluster_dims" in attrs: | ||
| # Extract cluster dimensions for SM90+ cluster launch | ||
| cluster_dims_attr = attrs["cluster_dims"] | ||
| cluster_dims = [int(cluster_dims_attr[i]) for i in range(len(cluster_dims_attr))] |
There was a problem hiding this comment.
Consider validating that cluster_dims has exactly 3 elements.
No length check is performed on the parsed cluster_dims_attr. If it has fewer than 3 elements, the downstream .format(…, *self.cluster_dims[function_name]) will raise an IndexError. Adding a quick assertion here provides a clear error message at the point of origin.
Proposed fix
if "cluster_dims" in attrs:
# Extract cluster dimensions for SM90+ cluster launch
cluster_dims_attr = attrs["cluster_dims"]
cluster_dims = [int(cluster_dims_attr[i]) for i in range(len(cluster_dims_attr))]
+ assert len(cluster_dims) == 3, (
+ f"cluster_dims must have exactly 3 elements, got {len(cluster_dims)}"
+ )📝 Committable suggestion
‼️ IMPORTANT
Carefully review the code before committing. Ensure that it accurately replaces the highlighted code, contains no missing lines, and has no issues with indentation. Thoroughly test & benchmark the code to ensure it meets the requirements.
| if "cluster_dims" in attrs: | |
| # Extract cluster dimensions for SM90+ cluster launch | |
| cluster_dims_attr = attrs["cluster_dims"] | |
| cluster_dims = [int(cluster_dims_attr[i]) for i in range(len(cluster_dims_attr))] | |
| if "cluster_dims" in attrs: | |
| # Extract cluster dimensions for SM90+ cluster launch | |
| cluster_dims_attr = attrs["cluster_dims"] | |
| cluster_dims = [int(cluster_dims_attr[i]) for i in range(len(cluster_dims_attr))] | |
| assert len(cluster_dims) == 3, ( | |
| f"cluster_dims must have exactly 3 elements, got {len(cluster_dims)}" | |
| ) |
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@tilelang/jit/adapter/wrapper.py` around lines 492 - 495, The parsed
cluster_dims_attr is not checked for length and may be shorter than 3 causing
downstream IndexError when used (see cluster_dims, cluster_dims_attr and
self.cluster_dims[function_name]); add a validation right after building
cluster_dims to ensure len(cluster_dims) == 3 and raise a clear ValueError or
AssertionError with a message referencing the offending attribute and
function_name (e.g., "expected 3 cluster_dims for function {function_name}, got
{len(cluster_dims)}") so the error is reported at the origin before any .format
or unpacking occurs.
There was a problem hiding this comment.
Actionable comments posted: 1
🧹 Nitpick comments (1)
testing/python/language/test_tilelang_language_cluster_launch.py (1)
55-57: Consider adding a tiny runtime smoke execution in this test path.Right now this validates compile/codegen patterns only. A minimal launch (small shapes) would catch runtime launch wiring regressions earlier.
🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@testing/python/language/test_tilelang_language_cluster_launch.py` around lines 55 - 57, Add a minimal runtime smoke execution to the test_cluster_launch path: after run_cython_cluster_launch() and run_tvm_ffi_cluster_launch(), invoke a tiny runtime invocation (e.g., run a single forward with very small shapes) to validate the actual launch/wiring instead of only codegen/compile. Either extend run_cython_cluster_launch and run_tvm_ffi_cluster_launch to accept a "smoke_run" flag that executes the compiled artifact with a small tensor, or add a new helper (e.g., smoke_execute_cluster or run_*_cluster_smoke) and call it from test_cluster_launch to run a tiny shape through the compiled module and assert success.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@testing/python/language/test_tilelang_language_cluster_launch.py`:
- Around line 33-50: The test run_tvm_ffi_cluster_launch is brittle because it
asserts an exact multiline snippet (check_str) against mod.get_host_source(),
which breaks on harmless formatting changes; update the test to normalize or
pattern-match the host source instead: fetch the string via
mod.get_host_source(), collapse or normalize whitespace (e.g., replace
consecutive whitespace/newlines with a single space) or use regex to assert the
presence of the essential tokens like "stack_ffi_any", "[3].type_index",
"[3].v_int64", "[4].v_int64", "[5].v_int64" and the numeric values 2,1,1 rather
than comparing the exact multiline layout; apply this change inside
run_tvm_ffi_cluster_launch replacing the check_str exact-match assertion with
the whitespace-normalized or regex-based assertions so the test passes despite
formatting changes.
---
Nitpick comments:
In `@testing/python/language/test_tilelang_language_cluster_launch.py`:
- Around line 55-57: Add a minimal runtime smoke execution to the
test_cluster_launch path: after run_cython_cluster_launch() and
run_tvm_ffi_cluster_launch(), invoke a tiny runtime invocation (e.g., run a
single forward with very small shapes) to validate the actual launch/wiring
instead of only codegen/compile. Either extend run_cython_cluster_launch and
run_tvm_ffi_cluster_launch to accept a "smoke_run" flag that executes the
compiled artifact with a small tensor, or add a new helper (e.g.,
smoke_execute_cluster or run_*_cluster_smoke) and call it from
test_cluster_launch to run a tiny shape through the compiled module and assert
success.
| assert 'clusterDim = {2, 1, 1}' in mod.get_host_source() | ||
|
|
||
|
|
||
| def run_tvm_ffi_cluster_launch(): | ||
| kernel = matmul(1024, 1024, 1024, 128, 128, 32) | ||
| mod = tilelang.compile(kernel, execution_backend="tvm_ffi") | ||
| check_str = r""" | ||
| (((TVMFFIAny*)stack_ffi_any)[3].type_index) = 1; | ||
| (((TVMFFIAny*)stack_ffi_any)[3].zero_padding) = 0; | ||
| (((TVMFFIAny*)stack_ffi_any)[3].v_int64) = ((int64_t)2); | ||
| (((TVMFFIAny*)stack_ffi_any)[4].type_index) = 1; | ||
| (((TVMFFIAny*)stack_ffi_any)[4].zero_padding) = 0; | ||
| (((TVMFFIAny*)stack_ffi_any)[4].v_int64) = ((int64_t)1); | ||
| (((TVMFFIAny*)stack_ffi_any)[5].type_index) = 1; | ||
| (((TVMFFIAny*)stack_ffi_any)[5].zero_padding) = 0; | ||
| (((TVMFFIAny*)stack_ffi_any)[5].v_int64) = ((int64_t)1); | ||
| """ | ||
| assert check_str in mod.get_host_source() |
There was a problem hiding this comment.
Make host-source assertions less formatting-fragile.
The current checks are tightly coupled to exact whitespace/layout, especially the multiline TVM FFI snippet, so harmless codegen formatting changes can fail the test.
Proposed refactor
+def _assert_tvm_ffi_cluster_dims(host_src: str) -> None:
+ required = (
+ "(((TVMFFIAny*)stack_ffi_any)[3].v_int64) = ((int64_t)2);",
+ "(((TVMFFIAny*)stack_ffi_any)[4].v_int64) = ((int64_t)1);",
+ "(((TVMFFIAny*)stack_ffi_any)[5].v_int64) = ((int64_t)1);",
+ )
+ for snippet in required:
+ assert snippet in host_src
+
def run_cython_cluster_launch():
kernel = matmul(1024, 1024, 1024, 128, 128, 32)
mod = tilelang.compile(kernel, execution_backend="cython")
- assert 'clusterDim = {2, 1, 1}' in mod.get_host_source()
+ host_src = mod.get_host_source()
+ assert "clusterDim" in host_src and "{2, 1, 1}" in host_src
@@
def run_tvm_ffi_cluster_launch():
kernel = matmul(1024, 1024, 1024, 128, 128, 32)
mod = tilelang.compile(kernel, execution_backend="tvm_ffi")
- check_str = r"""
- (((TVMFFIAny*)stack_ffi_any)[3].type_index) = 1;
- (((TVMFFIAny*)stack_ffi_any)[3].zero_padding) = 0;
- (((TVMFFIAny*)stack_ffi_any)[3].v_int64) = ((int64_t)2);
- (((TVMFFIAny*)stack_ffi_any)[4].type_index) = 1;
- (((TVMFFIAny*)stack_ffi_any)[4].zero_padding) = 0;
- (((TVMFFIAny*)stack_ffi_any)[4].v_int64) = ((int64_t)1);
- (((TVMFFIAny*)stack_ffi_any)[5].type_index) = 1;
- (((TVMFFIAny*)stack_ffi_any)[5].zero_padding) = 0;
- (((TVMFFIAny*)stack_ffi_any)[5].v_int64) = ((int64_t)1);
-"""
- assert check_str in mod.get_host_source()
+ _assert_tvm_ffi_cluster_dims(mod.get_host_source())🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed.
In `@testing/python/language/test_tilelang_language_cluster_launch.py` around
lines 33 - 50, The test run_tvm_ffi_cluster_launch is brittle because it asserts
an exact multiline snippet (check_str) against mod.get_host_source(), which
breaks on harmless formatting changes; update the test to normalize or
pattern-match the host source instead: fetch the string via
mod.get_host_source(), collapse or normalize whitespace (e.g., replace
consecutive whitespace/newlines with a single space) or use regex to assert the
presence of the essential tokens like "stack_ffi_any", "[3].type_index",
"[3].v_int64", "[4].v_int64", "[5].v_int64" and the numeric values 2,1,1 rather
than comparing the exact multiline layout; apply this change inside
run_tvm_ffi_cluster_launch replacing the check_str exact-match assertion with
the whitespace-normalized or regex-based assertions so the test passes despite
formatting changes.
There was a problem hiding this comment.
🧹 Nitpick comments (2)
src/target/rt_mod_cuda.cc (1)
44-48: Consider validatingcluster_dimsarray size.The code assumes
cluster_dimshas exactly 3 elements when pushing the X/Y/Z tags. If the array has fewer elements, downstream code accessingcluster_dims[0],[1],[2]could fail. Consider adding a size check for robustness.🛡️ Suggested defensive check
if (f->GetAttr<ffi::Array<Integer>>("cluster_dims").defined()) { + auto cluster_dims = f->GetAttr<ffi::Array<Integer>>("cluster_dims").value(); + ICHECK_EQ(cluster_dims.size(), 3) + << "cluster_dims attribute must have exactly 3 elements"; info.launch_param_tags.push_back(runtime::launch_param::kClusterDimX); info.launch_param_tags.push_back(runtime::launch_param::kClusterDimY); info.launch_param_tags.push_back(runtime::launch_param::kClusterDimZ); }🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@src/target/rt_mod_cuda.cc` around lines 44 - 48, The code unconditionally pushes kClusterDimX/Y/Z when f->GetAttr<ffi::Array<Integer>>("cluster_dims") is defined but does not verify the array length; modify the block that checks f->GetAttr<ffi::Array<Integer>>("cluster_dims") so you retrieve the attribute into a local (e.g., auto cluster_dims = f->GetAttr<ffi::Array<Integer>>("cluster_dims")), verify cluster_dims.defined() && cluster_dims.size() >= 3 before pushing runtime::launch_param::kClusterDimX, kClusterDimY, kClusterDimZ into info.launch_param_tags, and otherwise either skip pushing those tags or raise/log an appropriate error to avoid out-of-bounds access downstream.src/transform/lower_device_kernel_launch.cc (1)
107-122: Validatecluster_dimsarray size before element access.The code accesses
dims[0],dims[1],dims[2]without verifying thatdimshas exactly 3 elements. Ifcluster_dimshas fewer elements (e.g., due to a bug upstream), this will cause an out-of-bounds access.🛡️ Suggested defensive check
if (collector.info_.cluster_dims.defined()) { auto dims = collector.info_.cluster_dims.value(); + ICHECK_EQ(dims.size(), 3) + << "cluster_dims must have exactly 3 elements, got " << dims.size(); Array<PrimExpr> new_launch_args = {PrimExpr(dims[0]), PrimExpr(dims[1]), PrimExpr(dims[2])};🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@src/transform/lower_device_kernel_launch.cc` around lines 107 - 122, Validate the size of collector.info_.cluster_dims before indexing: check that collector.info_.cluster_dims.defined() and that dims (the value of collector.info_.cluster_dims) has at least 3 elements (e.g., dims.size() >= 3) before accessing dims[0], dims[1], dims[2]; if the size check fails, avoid the out-of-bounds access by logging or returning early and leaving collector.info_.launch_args/launch_params unchanged (or apply a safe fallback), and only build new_launch_args/new_launch_params with the cluster dim names (tvm::runtime::launch_param::kClusterDimX/Y/Z) when the check passes.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Nitpick comments:
In `@src/target/rt_mod_cuda.cc`:
- Around line 44-48: The code unconditionally pushes kClusterDimX/Y/Z when
f->GetAttr<ffi::Array<Integer>>("cluster_dims") is defined but does not verify
the array length; modify the block that checks
f->GetAttr<ffi::Array<Integer>>("cluster_dims") so you retrieve the attribute
into a local (e.g., auto cluster_dims =
f->GetAttr<ffi::Array<Integer>>("cluster_dims")), verify cluster_dims.defined()
&& cluster_dims.size() >= 3 before pushing runtime::launch_param::kClusterDimX,
kClusterDimY, kClusterDimZ into info.launch_param_tags, and otherwise either
skip pushing those tags or raise/log an appropriate error to avoid out-of-bounds
access downstream.
In `@src/transform/lower_device_kernel_launch.cc`:
- Around line 107-122: Validate the size of collector.info_.cluster_dims before
indexing: check that collector.info_.cluster_dims.defined() and that dims (the
value of collector.info_.cluster_dims) has at least 3 elements (e.g.,
dims.size() >= 3) before accessing dims[0], dims[1], dims[2]; if the size check
fails, avoid the out-of-bounds access by logging or returning early and leaving
collector.info_.launch_args/launch_params unchanged (or apply a safe fallback),
and only build new_launch_args/new_launch_params with the cluster dim names
(tvm::runtime::launch_param::kClusterDimX/Y/Z) when the check passes.
ℹ️ Review info
Configuration used: defaults
Review profile: CHILL
Plan: Pro
📒 Files selected for processing (9)
3rdparty/tvmexamples/gemm_sm100/gemm_tcgen5mma_ws_persistent.pysrc/target/rt_mod_cuda.ccsrc/transform/lower_device_kernel_launch.ccsrc/transform/lower_opaque_block.ccsrc/transform/split_host_device.cctesting/python/language/test_tilelang_language_cluster_launch.pytilelang/jit/adapter/wrapper.pytilelang/language/kernel.py
✅ Files skipped from review due to trivial changes (1)
- 3rdparty/tvm
- Introduced new built-in operations for cluster synchronization: `cluster_arrive_relaxed`, `cluster_arrive`, `cluster_wait`, `cluster_sync`, and `block_rank_in_cluster`. - Updated the CUDA code generator to handle these new operations. - Added corresponding Python bindings and documentation for the new cluster functions. - Included necessary header files for cluster operations in the CUDA code generation process.
- Introduced `ptx_arrive_cluster_barrier` built-in for arriving at cluster barriers. - Added `alloc_cluster_barrier` function for allocating cluster barrier buffers. - Updated CUDA code generation to handle the new cluster barrier operations. - Enhanced existing functions to support shared and cluster barrier scopes. - Added tests for cluster barrier functionality in the TileLang testing suite.
|
@regression-perf |
Performance Regression Test ReportTriggered by: @LeiWang1999 Results
Artifacts
|
This PR aims to support cluster launch, query, synchronization and barrier operations for CUDA sm90+, which is a preliminary for DSMEM and 2-cta tcgen5mma on Blackwell.
Summary by CodeRabbit
New Features
Tests
Refactor