Skip to content
Merged
Show file tree
Hide file tree
Changes from 11 commits
Commits
Show all changes
24 commits
Select commit Hold shift + click to select a range
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
24 changes: 24 additions & 0 deletions docs/tutorials/debug_tools_for_tilelang.md
Original file line number Diff line number Diff line change
Expand Up @@ -171,6 +171,30 @@ The output messages will include something like:
msg='hello world' BlockIdx=(0, 0, 0), ThreadIdx=(0, 0, 0): 0
```

### Visual Layout Inference For TileLang
The **Visual Layout Inference** tool automatically generates visual diagrams that illustrate the mapping between logical indices, thread IDs, and register file locations.

When TileLang performs layout inference, it determines how fragment buffers are distributed across threads. The visual layout tool captures this information and generates:
1. **Textual output**: A human-readable description of the layout mapping
2. **Visual diagrams**: Color-coded plots showing the thread-to-data mapping

The visual layout inference tool is controlled through the `TL_ENABLE_LAYOUT_VISUALIZATION` pass configuration. By default, visualization is **disabled** to avoid performance overhead during compilation.

`TL_ENABLE_LAYOUT_VISUALIZATION` accepts string values to control output formats:
- "True" or "all": Enabled, generates all formats (PDF, PNG, SVG)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
- "True" or "all": Enabled, generates all formats (PDF, PNG, SVG)
- "true" or "all": Enabled, generates all formats (PDF, PNG, SVG)

- "png": Generate PNG format only
- "pdf": Generate PDF format only
- "svg": Generate SVG format only

The output messages will include something like:
```
C_local layout inference:
Shape: [32, 32] -> [8]
Thread: _j // 16 * 64 + _i // 16 * 32 + _i % 8 * 4 + _j % 8 // 2
Index: [_j % 16 // 8 * 4 + _i % 16 // 8 * 2 + _j % 2]
```
Comment thread
coderabbitai[bot] marked this conversation as resolved.
Outdated


## Conclusion

By carefully examining intermediate representations (IR) before final code generation—and by leveraging runtime printing through `T.print`—one can quickly diagnose where index calculations, copy logic, or other kernel operations deviate from the intended behavior. This two-pronged approach (inspecting IR transformations and using runtime prints) is often sufficient for resolving generation and correctness issues in TileLang programs.
Expand Down
59 changes: 59 additions & 0 deletions examples/visual_layout_inference/visual_layout_inference.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
import tilelang
import tilelang.language as T

tilelang.disable_cache()


# use pass_configs to enable layout visualization
@tilelang.jit(
out_idx=[-1], pass_configs={tilelang.PassConfigKey.TL_ENABLE_LAYOUT_VISUALIZATION: "False"})
def matmul(M, N, K, block_M, block_N, block_K, dtype="float16", accum_dtype="float"):

@T.prim_func
def gemm(
A: T.Tensor((M, K), dtype),
B: T.Tensor((K, N), dtype),
C: T.Tensor((M, N), dtype),
):
with T.Kernel(T.ceildiv(N, block_N), T.ceildiv(M, block_M), threads=128) as (bx, by):
A_shared = T.alloc_shared((block_M, block_K), dtype)
B_shared = T.alloc_shared((block_K, block_N), dtype)
C_local = T.alloc_fragment((block_M, block_N), accum_dtype)

T.clear(C_local)
for k in T.Pipelined(T.ceildiv(K, block_K), num_stages=3):
T.copy(A[by * block_M, k * block_K], A_shared)
T.copy(B[k * block_K, bx * block_N], B_shared)
T.gemm(A_shared, B_shared, C_local)

T.copy(C_local, C[by * block_M, bx * block_N])

return gemm


def main():
kernel = matmul(128, 128, 128, 32, 32, 32)

import torch

a = torch.randn(128, 128).cuda().half()
b = torch.randn(128, 128).cuda().half()

c = kernel(a, b)

ref_c = a @ b

torch.testing.assert_close(c, ref_c, rtol=1e-2, atol=1e-2)
print("All check passed.")

# print the layout visualization result and save figures to ./tmp.
'''
C_local layout inference:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

would be better to be C_local inferenced layout:

Shape: [32, 32] -> [8]
Thread: _j // 16 * 64 + _i // 16 * 32 + _i % 8 * 4 + _j % 8 // 2
Index: [_j % 16 // 8 * 4 + _i % 16 // 8 * 2 + _j % 2]
'''


if __name__ == "__main__":
main()
2 changes: 2 additions & 0 deletions pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,8 @@ dependencies = [
# mldtypes should be greater than 0.5.1
# if you want to enable fp4
fp4 = ["ml-dtypes>=0.5.1"]
# if you want to enable layout inference visualization
vis = ["matplotlib"]

[build-system]
requires = ["cython>=3.0.0", "scikit-build-core"]
Expand Down
1 change: 1 addition & 0 deletions src/op/builtin.cc
Original file line number Diff line number Diff line change
Expand Up @@ -34,6 +34,7 @@ TVM_REGISTER_PASS_CONFIG_OPTION(kDisableVectorize256, Bool);
TVM_REGISTER_PASS_CONFIG_OPTION(kDisableWGMMA, Bool);
TVM_REGISTER_PASS_CONFIG_OPTION(kDisableShuffleElect, Bool);
TVM_REGISTER_PASS_CONFIG_OPTION(kStorageRewriteDetectInplace, Bool);
TVM_REGISTER_PASS_CONFIG_OPTION(kEnableLayoutVisualization, ffi::String);

DataType cuTensorMapType() { return DataType::UInt(8, 128); }

Expand Down
2 changes: 2 additions & 0 deletions src/op/builtin.h
Original file line number Diff line number Diff line change
Expand Up @@ -51,6 +51,8 @@ static constexpr const char *kDisableWGMMA = "tl.disable_wgmma";
static constexpr const char *kDisableShuffleElect = "tl.disable_shuffle_elect";
static constexpr const char *kStorageRewriteDetectInplace =
"tl.storage_rewrite_detect_inplace";
static constexpr const char *kEnableLayoutVisualization =
"tl.enable_layout_visualization";
/*!
* \brief Whether to disable dynamic tail split
*
Expand Down
1 change: 1 addition & 0 deletions tilelang/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,6 +137,7 @@ def _load_tile_lang_lib():
transform, # noqa: F401
language, # noqa: F401
engine, # noqa: F401
tools, # noqa: F401
)
from .autotuner import autotune # noqa: F401
from .transform import PassConfigKey # noqa: F401
Expand Down
4 changes: 4 additions & 0 deletions tilelang/engine/lower.py
Original file line number Diff line number Diff line change
Expand Up @@ -16,6 +16,7 @@
from tilelang.engine.param import KernelParam, CompiledArtifact
from tilelang.utils.target import determine_target
from tilelang.engine.phase import (
LayoutVisual,
PreLowerSemanticCheck,
LowerAndLegalize,
OptimizeForTarget,
Expand Down Expand Up @@ -249,6 +250,9 @@ def lower(
# Phase 1: Lower and legalize the IR
mod = LowerAndLegalize(mod, target)

# Visualize the layout
LayoutVisual(mod)

# Phase 2: Optimize the IR for the target
mod = OptimizeForTarget(mod, target)

Expand Down
20 changes: 20 additions & 0 deletions tilelang/engine/phase.py
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,26 @@ def should_force_let_inline(pass_ctx: PassContext | None = None) -> bool:
return bool(pass_ctx and pass_ctx.config.get(tilelang.PassConfigKey.TL_FORCE_LET_INLINE, False))


def should_enable_layout_visual(pass_ctx: PassContext | None = None) -> bool:
if pass_ctx is None:
pass_ctx = tilelang.transform.get_pass_context()

Comment on lines +70 to +75

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

⚠️ Potential issue | 🟠 Major

Harden should_enable_layout_visual against string-valued configs.

Right now enabled is returned directly from PassContext.config. If a user mistakenly sets "false" (string) instead of False (bool), if should_enable_layout_visual(): will still evaluate truthy and enable visualization.

A small defensive tweak keeps bool configs working while handling strings safely:

 def should_enable_layout_visual(pass_ctx: PassContext | None = None) -> bool:
     if pass_ctx is None:
         pass_ctx = tilelang.transform.get_pass_context()
-    enabled = pass_ctx.config.get(tilelang.PassConfigKey.TL_LAYOUT_VISUALIZATION_ENABLE, False)
-    return enabled
+    value = pass_ctx.config.get(tilelang.PassConfigKey.TL_LAYOUT_VISUALIZATION_ENABLE, False)
+    if isinstance(value, str):
+        v = value.strip().lower()
+        return bool(v and v != "false")
+    return bool(value)

This keeps the default “disabled unless explicitly enabled” behavior while avoiding surprises from accidental string usage.

🤖 Prompt for AI Agents
In tilelang/engine/phase.py around lines 70 to 75, the function returns the raw
value from pass_ctx.config which can be a string like "false" and still evaluate
truthy; change the logic to coerce and validate the config into a strict
boolean: fetch the raw value, if it is a bool return it, if it is a str
interpret only common true values (e.g. "true", "1", "yes") as True
(case-insensitive) and everything else as False, and default to False when the
key is missing or value is None.

config_value = pass_ctx.config.get(tilelang.PassConfigKey.TL_ENABLE_LAYOUT_VISUALIZATION.value,
"")

if config_value is None:
return False

config_str = str(config_value).strip().lower()
return bool(config_str and config_str != "false")


def LayoutVisual(mod: IRModule) -> None:
"""Apply layout visualization pass if enabled."""
if should_enable_layout_visual():
tilelang.tools.LayoutVisual()(mod)


def PreLowerSemanticCheck(mod: IRModule) -> None:
"""
Check whether the module is valid before lowering. If not, raise a user-friendly error
Expand Down
1 change: 1 addition & 0 deletions tilelang/tools/__init__.py
Original file line number Diff line number Diff line change
@@ -1,2 +1,3 @@
from .plot_layout import plot_layout # noqa: F401
from .Analyzer import *
from .layout_visual import LayoutVisual # noqa: F401
64 changes: 64 additions & 0 deletions tilelang/tools/layout_visual.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,64 @@
import tilelang
import tilelang.language as T
from tvm import tir
from tvm.tir import PyStmtExprVisitor

from tvm.tir.transform import prim_func_pass
from tilelang.tools.plot_layout import plot_layout

Comment thread
LeiWang1999 marked this conversation as resolved.
Outdated

def print_layout_format(layout: T.Fragment) -> str:

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This only works for fragment layout, I think we should rename it into print_fragment_format and do some type check there.

input_shape = layout.get_input_shape()
output_shape = layout.get_output_shape()
lines = [
f" Shape: {input_shape} -> {output_shape}", f" Thread: {layout.forward_thread}",
f" Index: {layout.forward_index}"
]

return "\n".join(lines)


@tir.functor.visitor
class _LayoutVisualVisitor(PyStmtExprVisitor):

def __init__(self, formats: str = "png"):
super().__init__()
self.layout_found = []
self.processed_layouts = set()
self.formats = formats

def visit_block_(self, op: tir.Block) -> None:
if "layout_map" in op.annotations:
layout_map = op.annotations["layout_map"]

for key, layout in layout_map.items():
if isinstance(layout, T.Fragment):
layout_id = str(layout)
if layout_id not in self.processed_layouts:
print(f"{key} layout inference:")
print(print_layout_format(layout))
plot_layout(layout, name=f"{key}_layout", formats=self.formats)
self.processed_layouts.add(layout_id)

self.visit_stmt(op.body)


def LayoutVisual():

def pass_fn(func: tir.PrimFunc, mod, ctx):
pass_ctx = tilelang.transform.get_pass_context()
config_value = pass_ctx.config.get(
tilelang.PassConfigKey.TL_ENABLE_LAYOUT_VISUALIZATION.value)

config_str = str(config_value).strip().lower()
if not config_str or config_str == "false":
return func
elif config_str == "true":
formats = "all"
else:
formats = config_str

_LayoutVisualVisitor(formats=formats).visit_stmt(func.body)
Comment thread
coderabbitai[bot] marked this conversation as resolved.
return func

return prim_func_pass(pass_fn, opt_level=0)
58 changes: 40 additions & 18 deletions tilelang/tools/plot_layout.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,13 @@
from __future__ import annotations
Comment thread
LeiWang1999 marked this conversation as resolved.
import tilelang.language as T


def plot_layout(layout: T.Layout,
def plot_layout(layout: T.Fragment,
save_directory="./tmp",
name: str = "layout",
colormap: str = "RdPu",
verbose: bool = False) -> None:
verbose: bool = False,
formats: str | list[str] = "png") -> None:
Comment thread
LeiWang1999 marked this conversation as resolved.
"""
Plot the layout of a buffer.

Expand All @@ -21,9 +23,10 @@ def plot_layout(layout: T.Layout,
The colormap to use for visualization (default is "RdPu").
verbose : bool, optional
If True, prints additional information about the mapping (default is False).

formats : str | list[str], optional
The formats to save the image in (default is "png").
Returns
-------
-------s

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Suggested change
-------s
-------

None
"""
import os
Expand Down Expand Up @@ -82,6 +85,12 @@ def plot_layout(layout: T.Layout,
raw_colors = [cmap(i) for i in range(num_threads)]
colors = raw_colors.copy()

# Show the distribution of registers in each thread of a warp.
warp_size = 32
spectral_camp = plt.get_cmap("hsv", warp_size * 6)
for i in range(warp_size):
colors[i] = spectral_camp(i * 6)
Comment on lines +88 to +101

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Guard warp coloring when thread count < warp size

The new warp-aware coloring loop assumes at least 32 threads, but colors is sized to num_threads; when a layout has fewer than 32 threads (common for small tiles), iterating for i in range(warp_size) writes past the end of the list and raises IndexError, so visualization fails before any plots are saved.

Useful? React with 👍 / 👎.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.


Comment thread
LeiWang1999 marked this conversation as resolved.
# Determine the number of rows and columns in the input shape
nrows, ncols = input_shape
# Adjust figure size to maintain square cells
Expand Down Expand Up @@ -191,17 +200,30 @@ def plot_layout(layout: T.Layout,
# Save the figure in multiple formats
plt.tight_layout()

# Save as PDF
pdf_path = tmp_directory / f"{name}.pdf"
plt.savefig(pdf_path, bbox_inches="tight")
print(f"Saved pdf format into {pdf_path}")

# Save as PNG
png_path = tmp_directory / f"{name}.png"
plt.savefig(png_path, bbox_inches="tight", transparent=False, dpi=255)
print(f"Saved png format into {png_path}")

# Save as SVG
svg_path = tmp_directory / f"{name}.svg"
plt.savefig(svg_path, bbox_inches="tight", format="svg")
print(f"Saved svg format into {svg_path}")
if isinstance(formats, str):
formats_str = formats.strip().lower()
if formats_str == 'all':
formats_list = ['pdf', 'png', 'svg']
elif "," in formats_str:
formats_list = [f.strip() for f in formats_str.split(',')]
else:
formats_list = [formats_str]
else:
raise TypeError(f"Expected str, but got {type(formats).__name__}. "
f"Please pass a string like 'png', 'pdf', 'svg', 'all', or 'png,pdf'.")
Comment on lines +218 to +222

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Accept documented list inputs for output formats

The formats parameter is annotated and documented as str | list[str], but the new parsing block immediately raises a TypeError for any non‑string input, so callers following the API and passing a list like ["png", "pdf"] cannot use the feature and the tool fails before plotting.

Useful? React with 👍 / 👎.


# Save the figure
if 'pdf' in formats_list:
pdf_path = tmp_directory / f"{name}.pdf"
plt.savefig(pdf_path, bbox_inches="tight")
print(f"Saved pdf format into {pdf_path}")

if 'png' in formats_list:
png_path = tmp_directory / f"{name}.png"
plt.savefig(png_path, bbox_inches="tight", transparent=False, dpi=255)
print(f"Saved png format into {png_path}")

if 'svg' in formats_list:
svg_path = tmp_directory / f"{name}.svg"
plt.savefig(svg_path, bbox_inches="tight", format="svg")
print(f"Saved svg format into {svg_path}")
8 changes: 8 additions & 0 deletions tilelang/transform/pass_config.py
Original file line number Diff line number Diff line change
Expand Up @@ -69,6 +69,14 @@ class PassConfigKey(str, Enum):
TL_FORCE_LET_INLINE = "tl.force_let_inline"
"""Force TileLang to inline let bindings during simplification. Default: False"""

TL_ENABLE_LAYOUT_VISUALIZATION = "tl.enable_layout_visualization"
"""Enable layout inference visualization. Accepts string values:
- "" or "false": disabled (default)
- "true" or "all": enabled, generate all formats (pdf, png, svg)
- "png", "pdf", "svg": enabled, generate specified format
- "png,svg": enabled, generate multiple formats (comma-separated)
""" ""

TL_STORAGE_REWRITE_DETECT_INPLACE = "tl.storage_rewrite_detect_inplace"
"""Control StorageRewrite inplace detection.

Expand Down
Loading