Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
The table of contents is too big for display.
Diff view
Diff view
  •  
  •  
  •  
Prev Previous commit
Next Next commit
Merge main into metal-gemm
  • Loading branch information
oraluben committed May 21, 2026
commit eaa9569fbb599e3fad905ac1426ef636a6ccff40
105 changes: 105 additions & 0 deletions .agents/skills/tilelang-tvm-ir/SKILL.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,105 @@
---
name: tilelang-tvm-ir
description: Use when editing TileLang C++ passes or TVM TIRX code that handles ObjectRef/NodeRef types such as For, Buffer, Var, SBlock, Stmt, PrimExpr, or their *Node raw node counterparts; especially when choosing function parameters, optional values, identity maps/sets, or equality checks.
---

# TileLang TVM IR Handle Conventions

## Core Rule

In TVM C++, `For`, `Buffer`, `Var`, `SBlock`, `Stmt`, `PrimExpr`, `SeqStmt`, etc. are `ObjectRef` smart handles. `ForNode`, `BufferNode`, `VarNode`, `SBlockNode`, etc. are raw node structs reached through visitor callbacks, `as<TNode>()`, `operator->`, or `.get()`.

When a value needs to cross a function boundary, be stored, be optional, be used as an identity key, or survive beyond a local inspection branch, prefer the handle type over `const *Node`.

## Preferred Patterns

- Function parameters and return values: use handles such as `For`, `Buffer`, `Var`, `SBlock`, `Stmt`, or `SeqStmt`.
- Nullable AST values: use `Optional<For>`, `Optional<SeqStmt>`, etc., not `const ForNode* = nullptr`.
- Identity maps and sets: use handle keys with TVM identity hashing:

```cpp
using BufferSet = std::unordered_set<Buffer, ObjectPtrHash, ObjectPtrEqual>;
using BufferMap = std::unordered_map<Buffer, Buffer, ObjectPtrHash, ObjectPtrEqual>;
using VarMap = std::unordered_map<Var, PrimExpr, ObjectPtrHash, ObjectPtrEqual>;
```

- Identity comparisons: use `.same_as(other)` when comparing two handles.
- Visitor callback node pointers: convert to a handle with `GetRef<T>(op)` when the value must be retained or passed elsewhere.

```cpp
Stmt VisitStmt_(const ForNode* op) final {
For loop = GetRef<For>(op);
Optional<For> candidate = FindPipelineLoop(loop->body);
if (candidate.defined() && candidate.value().same_as(loop)) {
...
}
}
```

- Pattern matching and local mutation may still use node pointers:
- `if (const auto* seq = stmt.as<SeqStmtNode>()) { ... }`
- `BufferStoreNode* n = store.CopyOnWrite();`
- visitor overrides such as `VisitStmt_(const SeqStmtNode* op)`

Keep these raw pointers local to the immediate inspection or mutation site.

## Avoid

- Passing `const ForNode*`, `const BufferNode*`, `const VarNode*`, or `const SBlockNode*` between helper functions when a handle exists.
- Storing raw node pointers in `std::unordered_map` or `std::unordered_set` for identity tracking.
- Using `.get()` as a key unless a callee requires a raw TVM node API and the pointer is not retained.
- Comparing handles through `.get() == other.get()`; prefer `.same_as()`.
- Reconstructing handles from raw pointers repeatedly when a handle is already available.

## Common Refactors

```cpp
// Before
const SeqStmtNode* pipeline_body_seq = nullptr;
pipeline_body_seq = seq_stmt;
ICHECK(pipeline_body_seq != nullptr);

// After
Optional<SeqStmt> pipeline_body_seq;
pipeline_body_seq = GetRef<SeqStmt>(seq_stmt);
ICHECK(pipeline_body_seq.defined());
SeqStmt pipeline_body = pipeline_body_seq.value();
```

```cpp
// Before
std::unordered_set<const BufferNode*> seen;
seen.insert(buffer.get());
if (seen.count(read->buffer.get())) { ... }

// After
BufferSet seen;
seen.insert(buffer);
if (seen.count(read->buffer)) { ... }
```

```cpp
// Before
std::unordered_set<const VarNode*> vars;
vars.insert(loop->loop_var.get());
bool uses = UsesVar(expr, [&](const VarNode* vn) {
return vars.count(vn) > 0;
});

// After
VarSet vars;
vars.insert(loop->loop_var);
bool uses = UsesVar(expr, [&](const VarNode* vn) {
return vars.count(GetRef<Var>(vn)) > 0;
});
```

## Review Checklist

When reviewing TileLang TIR passes, search for:

```bash
rg -n "std::unordered_(set|map)<const .*Node \\*|const (For|SeqStmt|SBlock).*Node \\*|\\.get\\(\\) ==|\\.find\\([^\\n]*\\.get\\(\\)|\\.count\\([^\\n]*\\.get\\(\\)|\\.insert\\([^\\n]*\\.get\\(\\)" src/transform
```

Do not mechanically remove every raw node pointer. Keep visitor signatures, `as<TNode>()` pattern checks, and `CopyOnWrite()` mutation pointers. Refactor only the places that store, pass, compare, or key identities through raw pointers.
6 changes: 3 additions & 3 deletions .github/workflows/ci.yml
Original file line number Diff line number Diff line change
Expand Up @@ -44,19 +44,19 @@ jobs:
fetch-depth: 0
submodules: recursive

- name: Setup Python 3.9
- name: Setup Python 3.10
id: setup-pylowest
uses: actions/setup-python@v6
with:
python-version: "3.9"
python-version: "3.10"
update-environment: true
cache: pip
cache-dependency-path: |
pyproject.toml
requirements*.txt
.pre-commit-config.yaml

- name: Check AST with Python 3.9
- name: Check AST with Python 3.10
run: |
"${{ steps.setup-pylowest.outputs.python-path }}" -m compileall -q -f tilelang

Expand Down
12 changes: 5 additions & 7 deletions .github/workflows/dist.yml
Original file line number Diff line number Diff line change
Expand Up @@ -18,8 +18,7 @@ on:
- CMakeLists.txt
- version_provider.py
- .github/workflows/dist.yml
# temporarily add to dist check
# until we have type checking in ci / move to python 3.10
# Type aliases can affect package import/build behavior.
- tilelang/_typing.py
release:
types:
Expand Down Expand Up @@ -115,12 +114,11 @@ jobs:
strategy:
matrix:
target:
# Build wheels for different Python ABIs.
# Windows CUDA 13.0 uses cp310 because PyTorch cu130 does not publish cp39 wheels.
- { runner: ubuntu-latest, toolkit: "CUDA-12.8", test_backends: "cu118 cu130", python_version: "3.9" }
- { runner: ubuntu-24.04-arm, toolkit: "CUDA-12.8", test_backends: "cu126 cu130", python_version: "3.9" }
# Build wheels for the minimum supported Python ABI.
- { runner: ubuntu-latest, toolkit: "CUDA-12.8", test_backends: "cu118 cu130", python_version: "3.10" }
- { runner: ubuntu-24.04-arm, toolkit: "CUDA-12.8", test_backends: "cu126 cu130", python_version: "3.10" }
- { runner: windows-latest, toolkit: "CUDA-13.0", test_backends: "cu130", python_version: "3.10" }
- { runner: macos-latest, toolkit: "Metal", python_version: "3.9" }
- { runner: macos-latest, toolkit: "Metal", python_version: "3.10" }
# - "3.14t" # let user to build from source for now
# TODO: Add cp315-abi3.abi3t after PEP 803
fail-fast: false
Expand Down
Loading
Loading
You are viewing a condensed version of this merge commit. You can view the full changes here.