Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
16 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
18 changes: 16 additions & 2 deletions docs/arch/tvmscript.rst
Original file line number Diff line number Diff line change
Expand Up @@ -45,6 +45,20 @@ function decorators retain definitions until module construction, which can decl
signatures before building bodies. Captured Python values belong to the source
context; symbolic IR values are created and resolved by builders.

Annotations read concrete values from their definition scope, preserving missing-name
errors when a value is used. Create external symbols with ``n = I.dynamic("n")``
or the identical ``T.dynamic`` and ``Ts.dynamic`` constructors. Each call creates a
fresh native variable, defaulting to int64. On Python 3.12+, explicit headers such as
``def f[n, k: T.int32](...)`` declare local symbols; ``n: int`` retains the int64
default. Quote the whole annotation or use ``from __future__ import annotations``
to defer eager Python evaluation of header symbols. Captured runtime ``typing.TypeVar``
objects are not script symbols; ordinary Python typing uses remain unaffected.

An explicit scalar annotation ``n: n`` preserves a captured native symbol's identity.
An independently typed parameter such as ``n: T.int32`` and ordinary body locals
shadow definition captures normally. Annotation classes, Python unions and deferred
return-constructor evaluation retain their builder behavior.

Syntax and construction protocol
--------------------------------

Expand All @@ -55,8 +69,8 @@ builder frames. The namespace owns the meaning of these operations and the suppo
IR constructs.

``tvm.script.parser.protocol_registry`` records syntax policies under registered
namespace paths. These policies identify declarations and arguments that need special
handling, such as symbolic expression strings. Source aliases resolve to those paths;
namespace paths. These policies identify scalar annotations, mutable declarations and
result span handling. Symbolic shapes use concrete expressions. Source aliases resolve to those paths;
ordinary Python calls remain calls in the generated program. Explicit ``constexpr``
markers select host control flow during construction.

Expand Down
27 changes: 16 additions & 11 deletions docs/deep_dive/relax/tutorials/relax_creation.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,17 +43,19 @@
from tvm.script import s_tir as Ts
from tvm.script import tirx as T

n = T.dynamic("n")


@I.ir_module
class RelaxModule:
@R.function
def forward(
data: R.Tensor(("n", 784), dtype="float32"),
data: R.Tensor((n, 784), dtype="float32"),
w0: R.Tensor((128, 784), dtype="float32"),
b0: R.Tensor((128,), dtype="float32"),
w1: R.Tensor((10, 128), dtype="float32"),
b1: R.Tensor((10,), dtype="float32"),
) -> R.Tensor(("n", 10), dtype="float32"):
) -> R.Tensor((n, 10), dtype="float32"):
with R.dataflow():
lv0 = R.matmul(data, R.permute_dims(w0)) + b0
lv1 = R.nn.relu(lv0)
Expand All @@ -70,12 +72,14 @@ def forward(
# TensorIR functions in Relax function.


n = T.dynamic("n", "int64")
m = T.dynamic("m", "int64")


@I.ir_module
class RelaxModuleWithTIR:
@Ts.prim_func
def relu(x: T.handle, y: T.handle):
n = T.int64()
m = T.int64()
X = T.match_buffer(x, (n, m), "float32")
Y = T.match_buffer(y, (n, m), "float32")
for i, j in T.grid(n, m):
Expand All @@ -85,13 +89,12 @@ def relu(x: T.handle, y: T.handle):

@R.function
def forward(
data: R.Tensor(("n", 784), dtype="float32"),
data: R.Tensor((n, 784), dtype="float32"),
w0: R.Tensor((128, 784), dtype="float32"),
b0: R.Tensor((128,), dtype="float32"),
w1: R.Tensor((10, 128), dtype="float32"),
b1: R.Tensor((10,), dtype="float32"),
) -> R.Tensor(("n", 10), dtype="float32"):
n = T.int64()
) -> R.Tensor((n, 10), dtype="float32"):
cls = RelaxModuleWithTIR
with R.dataflow():
lv0 = R.matmul(data, R.permute_dims(w0)) + b0
Expand Down Expand Up @@ -165,11 +168,13 @@ def forward(self, x):
# Tensor Expression(TE), TensorIR functions or other TVM packed functions.


M = T.dynamic("M", "int64")
N = T.dynamic("N", "int64")
K = T.dynamic("K", "int64")


@Ts.prim_func
def tir_linear(x: T.handle, w: T.handle, b: T.handle, z: T.handle):
M = T.int64()
N = T.int64()
K = T.int64()
X = T.match_buffer(x, (M, K), "float32")
W = T.match_buffer(w, (N, K), "float32")
B = T.match_buffer(b, (N,), "float32")
Expand Down Expand Up @@ -227,7 +232,7 @@ def forward(self, x):
# customized pass.

bb = relax.BlockBuilder()
n = T.int64()
n = T.dynamic("n", "int64")
x = relax.Var("x", R.Tensor((n, 784), "float32"))
fc1_weight = relax.Var("fc1_weight", R.Tensor((128, 784), "float32"))
fc1_bias = relax.Var("fc1_bias", R.Tensor((128,), "float32"))
Expand Down
11 changes: 6 additions & 5 deletions docs/deep_dive/tensor_ir/tutorials/tir_creation.py
Original file line number Diff line number Diff line change
Expand Up @@ -178,15 +178,16 @@ def mm_relu(
# be used to ascertain the shape and data type of a TensorIR.


# Dynamic shape definition
M = T.dynamic("M", "int32")
N = T.dynamic("N", "int32")
K = T.dynamic("K", "int32")


@I.ir_module
class DynamicShapeModule:
@Ts.prim_func
def mm_relu(a: T.handle, b: T.handle, c: T.handle):
# Dynamic shape definition
M = T.int32()
N = T.int32()
K = T.int32()

# Bind the input buffers with the dynamic shapes
A = T.match_buffer(a, [M, K], dtype)
B = T.match_buffer(b, [K, N], dtype)
Expand Down
12 changes: 6 additions & 6 deletions docs/how_to/tutorials/mix_python_and_tvm_with_pymodule.py
Original file line number Diff line number Diff line change
Expand Up @@ -129,12 +129,12 @@ def forward(self, x, y):
# immediately — no recompilation needed.

if RUN_EXAMPLE:
n = T.dynamic("n", "int32")

@R.py_module
class DebugModule(BasePyModule):
@Ts.prim_func
def matmul_tir(var_A: T.handle, var_B: T.handle, var_C: T.handle):
n = T.int32()
A = T.match_buffer(var_A, (n, 4), "float32")
B = T.match_buffer(var_B, (4, 3), "float32")
C = T.match_buffer(var_C, (n, 3), "float32")
Expand Down Expand Up @@ -399,22 +399,22 @@ def main(
# tensors at call time, so the same module handles different sizes without recompilation.

if RUN_EXAMPLE:
n = T.dynamic("n", "int64")

@R.py_module
class DynamicModule(BasePyModule):
@Ts.prim_func
def scale_tir(var_x: T.handle, var_out: T.handle):
n = T.int64()
x = T.match_buffer(var_x, (n,), "float32")
out = T.match_buffer(var_out, (n,), "float32")
for i in T.serial(n):
out[i] = x[i] * T.float32(2.0)

@R.function
def add_relax(
x: R.Tensor(("n",), "float32"),
y: R.Tensor(("n",), "float32"),
) -> R.Tensor(("n",), "float32"):
x: R.Tensor((n,), "float32"),
y: R.Tensor((n,), "float32"),
) -> R.Tensor((n,), "float32"):
return R.add(x, y)

mod = DynamicModule(device=tvm.cpu(0), target="llvm")
Expand All @@ -435,7 +435,7 @@ def add_relax(
print("add_relax(len=10):", out10)

# Python → TIR with symbolic output shape
n = T.int64()
n = T.dynamic("n", "int64")
x7 = torch.randn(7)
scaled = mod.call_tir("scale_tir", [x7], relax.TensorType((n,), "float32"))
print("scale_tir(len=7):", scaled)
Expand Down
2 changes: 1 addition & 1 deletion docs/reference/api/python/script/parser.rst
Original file line number Diff line number Diff line change
Expand Up @@ -32,7 +32,7 @@ registered through :func:`tvm.script.parser.register_namespace` and
:func:`tvm.script.parser.register_namespace_initializer`.

.. automodule:: tvm.script.parser.protocol_registry
:members: constexpr, args_policy, register_type_var_decl, mutable_cell_decl, result_span, module_decorator, declaration_kind
:members: constexpr, register_scalar_annotation, mutable_cell_decl, result_span, module_decorator, declaration_kind

The language variant aliases below share the public construction namespaces documented
in :doc:`script`. Parser entry points above use the canonical frontend.
Expand Down
4 changes: 3 additions & 1 deletion jvm/core/src/test/scripts/prepare_test_libs.py
Original file line number Diff line number Diff line change
Expand Up @@ -21,16 +21,18 @@

import tvm
from tvm import relax, te
from tvm.script import ir as I
from tvm.script import relax as R


def prepare_relax_lib(base_path):
pipeline = relax.get_pipeline()
n = I.dynamic("n")

@tvm.script.ir_module
class Mod:
@R.function
def main(x: R.Tensor(["n"], "float32"), y: R.Tensor(["n"], "float32")):
def main(x: R.Tensor([n], "float32"), y: R.Tensor([n], "float32")):
lv0 = R.add(x, y)
return lv0

Expand Down
9 changes: 7 additions & 2 deletions python/tvm/relax/backend/gpu_generic/cumsum.py
Original file line number Diff line number Diff line change
Expand Up @@ -174,10 +174,12 @@ def update_cross_block(
bx > 0, source[by, src_offset + bx - 1], 0
)

m = T.dynamic("m")
n = T.dynamic("n")

@Ts.prim_func(private=True)
def cumsum(var_a: T.handle, var_out: T.handle):
T.func_attr({"tirx.is_scheduled": True}) # prevent further scheduling
m, n = T.int64(), T.int64()
A = T.match_buffer(var_a, [m, n], dtype=in_dtype)
Out = T.match_buffer(var_out, [m, n], dtype=out_dtype)
Tmp = T.alloc_buffer([m, n], dtype=out_dtype)
Expand Down Expand Up @@ -253,10 +255,13 @@ def gpu_3d_axis_1_cumsum(
out_dtype = out_dtype or in_dtype
TX = T.int64(tx_len)

outer = T.dynamic("outer")
scan = T.dynamic("scan")
inner = T.dynamic("inner")

@Ts.prim_func(private=True)
def cumsum(var_a: T.handle, var_out: T.handle):
T.func_attr({"tirx.is_scheduled": True})
outer, scan, inner = T.int64(), T.int64(), T.int64()
A = T.match_buffer(var_a, [outer, scan, inner], dtype=in_dtype)
Out = T.match_buffer(var_out, [outer, scan, inner], dtype=out_dtype)

Expand Down
11 changes: 8 additions & 3 deletions python/tvm/relax/backend/gpu_generic/sampling.py
Original file line number Diff line number Diff line change
Expand Up @@ -259,6 +259,10 @@ def single_batch_sampling(

aggregate[()] += step_aggregate[()]

n = T.dynamic("n")
vocab_size = T.dynamic("vocab_size")
batch_size = T.dynamic("batch_size")

@Ts.prim_func
def parallel_sampling_from_prob(
var_prob: T.handle,
Expand All @@ -267,7 +271,6 @@ def parallel_sampling_from_prob(
var_sampled_token_ids: T.handle,
):
T.func_attr({"tirx.is_scheduled": True})
n, vocab_size, batch_size = T.int64(), T.int64(), T.int64()
# match buffers
prob = T.match_buffer(var_prob, (n, vocab_size), prob_dtype)
uniform_samples = T.match_buffer(var_uniform_samples, (batch_size, 1), sample_dtype)
Expand Down Expand Up @@ -318,11 +321,13 @@ def generic_get_sample_index(
):
"""Generate a generic get_sample_index kernel."""

batch = T.dynamic("batch")
vocab_size = T.dynamic("vocab_size")
out_batch = T.dynamic("out_batch")

@Ts.prim_func(private=True)
def _get_sample_index(A: T.handle, B: T.handle, C: T.handle, D: T.handle):
batch, vocab_size = T.int64(), T.int64()
prob = T.match_buffer(A, (batch, vocab_size), prob_dtype)
out_batch = T.int64()
usample = T.match_buffer(B, (out_batch, 1), sample_dtype)
sample_indices = T.match_buffer(C, (out_batch, 1), sample_indices_dtype)
output_index = T.match_buffer(D, (out_batch, 1), dtype)
Expand Down
10 changes: 7 additions & 3 deletions python/tvm/relax/block_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -470,15 +470,16 @@ def te_func(args, args_dict, msg):

.. code-block:: python

m = T.dynamic("m")
n = T.dynamic("n")

@tvm.script.ir_module
class Module:
@Ts.prim_func
def te_func(var_rxplaceholder: T.handle, var_rxplaceholder_1: T.handle,
var_compute: T.handle) -> None:
# function attr dict
T.func_attr({"tirx.noalias": True})
m = T.int64()
n = T.int64()
rxplaceholder = T.match_buffer(var_rxplaceholder, [n, m], dtype="float32")
rxplaceholder_1 = T.match_buffer(var_rxplaceholder_1, [n, m], dtype="float32")
compute = T.match_buffer(var_compute, [128, 128], dtype="float32")
Expand Down Expand Up @@ -519,6 +520,9 @@ def te_func(A):

.. code-block:: python

m = T.dynamic("m")
n = T.dynamic("n")

@tvm.script.ir_module
class Module:
@Ts.prim_func
Expand Down Expand Up @@ -642,7 +646,7 @@ def emit_func_output(
# `bb.function()`, then any variables provided from the params
# are not in scope. Otherwise, TIR variables used in dynamic
# inputs are removed as undefined (e.g. Replacing
# `R.Tensor(["batch_size"])` with `R.Tensor(ndims=1)`).
# `R.Tensor([batch_size])` with `R.Tensor(ndims=1)`).
self.begin_scope(self._func._params)
try:
seqe = self.normalize(seqe)
Expand Down
Loading
Loading