Release v4.0.0 (#2294)

This commit is contained in:
Kihiro Bando
2025-05-13 15:55:29 -04:00
committed by GitHub
parent ad7b2f5e84
commit f115c3f854
299 changed files with 51495 additions and 4413 deletions

View File

@@ -0,0 +1,31 @@
# Copyright
Copyright (c) 2025 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
SPDX-License-Identifier: BSD-3-Clause
```
Redistribution and use in source and binary forms, with or without
modification, are permitted provided that the following conditions are met:
1. Redistributions of source code must retain the above copyright notice, this
list of conditions and the following disclaimer.
2. Redistributions in binary form must reproduce the above copyright notice,
this list of conditions and the following disclaimer in the documentation
and/or other materials provided with the distribution.
3. Neither the name of the copyright holder nor the names of its
contributors may be used to endorse or promote products derived from
this software without specific prior written permission.
THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS"
AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE
IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE
DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE
FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL
DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR
SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER
CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY,
OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE
OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE.
```

View File

@@ -0,0 +1,648 @@
{
"cells": [
{
"cell_type": "markdown",
"id": "0e95f0df-4d1a-4e2e-92ff-90539bb4c517",
"metadata": {},
"source": [
"# Example 06: CUDA Graphs\n",
"\n",
"In this example we demonstrate how to use CUDA graphs through PyTorch with CuTe DSL.\n",
"The process of interacting with PyTorch's CUDA graph implementation requires exposing PyTorch's CUDA streams to CUTLASS.\n",
"\n",
"To use CUDA graphs with Blackwell requires a version of PyTorch that supports Blackwell.\n",
"This can be obtained through:\n",
"- The [PyTorch NGC](https://catalog.ngc.nvidia.com/orgs/nvidia/containers/pytorch)\n",
"- [PyTorch 2.7 with CUDA 12.8 or later](https://pytorch.org/) (e.g., `pip3 install torch torchvision torchaudio --index-url https://download.pytorch.org/whl/cu128`)\n",
"- Building PyTorch directly with your version of CUDA."
]
},
{
"cell_type": "code",
"execution_count": 1,
"id": "46b8fb6f-9ac5-4a3d-b765-b6476f182bf7",
"metadata": {},
"outputs": [],
"source": [
"# import torch for CUDA graphs\n",
"import torch\n",
"import cutlass\n",
"import cutlass.cute as cute\n",
"# import CUstream type from the cuda driver bindings\n",
"from cuda.bindings.driver import CUstream\n",
"# import the current_stream function from torch\n",
"from torch.cuda import current_stream"
]
},
{
"cell_type": "markdown",
"id": "bcf5e06e-1f5b-4d72-ad73-9b36efb78ca0",
"metadata": {},
"source": [
"## Kernel Creation\n",
"\n",
"We create a kernel which prints \"Hello world\" as well as a host function to launch the kernel.\n",
"We then compile the kernel for use in our graph, by passing in a default stream.\n",
"\n",
"Kernel compilation before graph capture is required since CUDA graphs cannot JIT compile kernels during graph execution."
]
},
{
"cell_type": "code",
"execution_count": 2,
"id": "0c2a6ca8-98d7-4837-b91f-af769ca8fcd8",
"metadata": {},
"outputs": [],
"source": [
"@cute.kernel\n",
"def hello_world_kernel():\n",
" \"\"\"\n",
" A kernel that prints hello world\n",
" \"\"\"\n",
" cute.printf(\"Hello world\")\n",
"\n",
"@cute.jit\n",
"def hello_world(stream : CUstream):\n",
" \"\"\"\n",
" Host function that launches our (1,1,1), (1,1,1) grid in stream\n",
" \"\"\"\n",
" hello_world_kernel().launch(grid=[1, 1, 1], block=[1, 1, 1], stream=stream)\n",
"\n",
"# Grab a stream from PyTorch, this will also initialize our context\n",
"# so we can omit cutlass.cuda.initialize_cuda_context()\n",
"stream = current_stream()\n",
"hello_world_compiled = cute.compile(hello_world, CUstream(stream.cuda_stream))"
]
},
{
"cell_type": "markdown",
"id": "ecc850af-09f8-4a29-9c93-ff31fbb9326f",
"metadata": {},
"source": [
"## Creating and replaying a CUDA Graph\n",
"\n",
"We create a stream through torch as well as a graph.\n",
"When we create the graph we can pass the stream we want to capture to torch. We similarly run the compiled kernel with the stream passed as a CUstream.\n",
"\n",
"Finally we can replay our graph and synchronize."
]
},
{
"cell_type": "code",
"execution_count": 3,
"id": "f673e5ae-42bb-44d0-b652-3280606181c4",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Hello world\n",
"Hello world\n"
]
}
],
"source": [
"# Create a CUDA Graph\n",
"g = torch.cuda.CUDAGraph()\n",
"# Capture our graph\n",
"with torch.cuda.graph(g):\n",
" # Turn our torch Stream into a cuStream stream.\n",
" # This is done by getting the underlying CUstream with .cuda_stream\n",
" graph_stream = CUstream(current_stream().cuda_stream)\n",
" # Run 2 iterations of our compiled kernel\n",
" for _ in range(2):\n",
" # Run our kernel in the stream\n",
" hello_world_compiled(graph_stream)\n",
"\n",
"# Replay our graph\n",
"g.replay()\n",
"# Synchronize all streams (equivalent to cudaDeviceSynchronize() in C++)\n",
"torch.cuda.synchronize()"
]
},
{
"attachments": {},
"cell_type": "markdown",
"id": "db76d9c3-7617-4bf2-b326-11982e6803bf",
"metadata": {},
"source": [
"Our run results in the following execution when viewed in NSight Systems:\n",
"\n",
"![Image of two hello world kernels run back to back in a CUDA graph](images/cuda_graphs_image.png)\n",
"\n",
"We can observe the launch of the two kernels followed by a `cudaDeviceSynchronize()`.\n",
"\n",
"Now we can confirm that this minimizes some launch overhead:"
]
},
{
"cell_type": "code",
"execution_count": 4,
"id": "3ebe15bf-dc97-42e9-913c-224ecfb472e8",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n",
"Hello world\n"
]
}
],
"source": [
"# Get our CUDA stream from PyTorch\n",
"stream = CUstream(current_stream().cuda_stream)\n",
"\n",
"# Create a larger CUDA Graph of 100 iterations\n",
"g = torch.cuda.CUDAGraph()\n",
"# Capture our graph\n",
"with torch.cuda.graph(g):\n",
" # Turn our torch Stream into a cuStream stream.\n",
" # This is done by getting the underlying CUstream with .cuda_stream\n",
" graph_stream = CUstream(current_stream().cuda_stream)\n",
" # Run 2 iterations of our compiled kernel\n",
" for _ in range(100):\n",
" # Run our kernel in the stream\n",
" hello_world_compiled(graph_stream)\n",
"\n",
"# Create CUDA events for measuring performance\n",
"start = torch.cuda.Event(enable_timing=True)\n",
"end = torch.cuda.Event(enable_timing=True)\n",
"\n",
"# Run our kernel to warm up the GPU\n",
"for _ in range(100):\n",
" hello_world_compiled(stream)\n",
"\n",
"# Record our start time\n",
"start.record()\n",
"# Run 100 kernels\n",
"for _ in range(100):\n",
" hello_world_compiled(stream)\n",
"# Record our end time\n",
"end.record()\n",
"# Synchronize (cudaDeviceSynchronize())\n",
"torch.cuda.synchronize()\n",
"\n",
"# Calculate the time spent when launching kernels in a stream\n",
"# Results are in ms\n",
"stream_time = start.elapsed_time(end) \n",
"\n",
"# Warmup our GPU again\n",
"g.replay()\n",
"# Record our start time\n",
"start.record()\n",
"# Run our graph\n",
"g.replay()\n",
"# Record our end time\n",
"end.record()\n",
"# Synchronize (cudaDeviceSynchronize())\n",
"torch.cuda.synchronize()\n",
"\n",
"# Calculate the time spent when launching kernels in a graph\n",
"# units are ms\n",
"graph_time = start.elapsed_time(end)"
]
},
{
"cell_type": "code",
"execution_count": 5,
"id": "12b8151a-46b3-4c99-9945-301f6b628131",
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"8.94% speedup when using CUDA graphs for this kernel!\n"
]
}
],
"source": [
"# Print out speedup when using CUDA graphs\n",
"percent_speedup = (stream_time - graph_time) / graph_time\n",
"print(f\"{percent_speedup * 100.0:.2f}% speedup when using CUDA graphs for this kernel!\")"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.5"
}
},
"nbformat": 4,
"nbformat_minor": 5
}

File diff suppressed because it is too large Load Diff

View File

@@ -0,0 +1,310 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [],
"source": [
"from typing import List\n",
"\n",
"import cutlass\n",
"import cutlass.cute as cute"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Understanding data structure in CuTe DSL\n",
"\n",
"In most cases, data structures in CuTe DSL work the same as Python data structures with the notable difference that Python data structures in most cases are considered as static data which are interpreted by the DSL compiler embedded inside Python interpreter.\n",
"\n",
"To differentiate between compile-time and runtime values, CuTe DSL introduces primitive types that \n",
"represent dynamic values in JIT-compiled code.\n",
"\n",
"CuTe DSL provides a comprehensive set of primitive numeric types for representing dynamic values at \n",
"runtime. These types are formally defined within the CuTe DSL typing system:\n",
"\n",
"### Integer Types\n",
"- `Int8` - 8-bit signed integer\n",
"- `Int16` - 16-bit signed integer \n",
"- `Int32` - 32-bit signed integer\n",
"- `Int64` - 64-bit signed integer\n",
"- `Int128` - 128-bit signed integer\n",
"- `Uint8` - 8-bit unsigned integer\n",
"- `Uint16` - 16-bit unsigned integer\n",
"- `Uint32` - 32-bit unsigned integer\n",
"- `Uint64` - 64-bit unsigned integer\n",
"- `Uint128` - 128-bit unsigned integer\n",
"\n",
"### Floating Point Types\n",
"- `Float16` - 16-bit floating point\n",
"- `Float32` - 32-bit floating point \n",
"- `Float64` - 64-bit floating point\n",
"- `BFloat16` - Brain Floating Point format (16-bit)\n",
"- `TFloat32` - Tensor Float32 format (reduced precision format used in tensor operations)\n",
"- `Float8E4M3` - 8-bit floating point with 4-bit exponent and 3-bit mantissa\n",
"- `Float8E5M2` - 8-bit floating point with 5-bit exponent and 2-bit mantissa\n",
"\n",
"These specialized types are designed to represent dynamic values in CuTe DSL code that will be \n",
"evaluated at runtime, in contrast to Python's built-in numeric types which are evaluated during \n",
"compilation.\n",
"\n",
"### Example usage:\n",
"\n",
"```python\n",
"x = cutlass.Int32(5) # Creates a 32-bit integer\n",
"y = cutlass.Float32(3.14) # Creates a 32-bit float\n",
"\n",
"@cute.jit\n",
"def foo(a: cutlass.Int32): # annotate `a` as 32-bit integer passed to jit function via ABI\n",
" ...\n",
"```\n",
"To differentiate between compile-time and runtime values, CuTe DSL introduces primitive types that \n",
"represent dynamic values in JIT-compiled code.\n",
"\n",
"CuTe DSL provides a comprehensive set of primitive numeric types for representing dynamic values at \n",
"runtime. These types are formally defined within the CuTe DSL typing system:\n",
"\n",
"### Integer Types\n",
"- `Int8` - 8-bit signed integer\n",
"- `Int16` - 16-bit signed integer \n",
"- `Int32` - 32-bit signed integer\n",
"- `Int64` - 64-bit signed integer\n",
"- `Int128` - 128-bit signed integer\n",
"- `Uint8` - 8-bit unsigned integer\n",
"- `Uint16` - 16-bit unsigned integer\n",
"- `Uint32` - 32-bit unsigned integer\n",
"- `Uint64` - 64-bit unsigned integer\n",
"- `Uint128` - 128-bit unsigned integer\n",
"\n",
"### Floating Point Types\n",
"- `Float16` - 16-bit floating point\n",
"- `Float32` - 32-bit floating point \n",
"- `Float64` - 64-bit floating point\n",
"- `BFloat16` - Brain Floating Point format (16-bit)\n",
"- `TFloat32` - Tensor Float32 format (reduced precision format used in tensor operations)\n",
"- `Float8E4M3` - 8-bit floating point with 4-bit exponent and 3-bit mantissa\n",
"- `Float8E5M2` - 8-bit floating point with 5-bit exponent and 2-bit mantissa\n",
"\n",
"These specialized types are designed to represent dynamic values in CuTe DSL code that will be \n",
"evaluated at runtime, in contrast to Python's built-in numeric types which are evaluated during \n",
"compilation.\n",
"\n",
"### Example usage:\n",
"\n",
"```python\n",
"x = cutlass.Int32(5) # Creates a 32-bit integer\n",
"y = cutlass.Float32(3.14) # Creates a 32-bit float\n",
"\n",
"@cute.jit\n",
"def foo(a: cutlass.Int32): # annotate `a` as 32-bit integer passed to jit function via ABI\n",
" ...\n",
"```"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"a(static) = ?\n",
"b(static) = ?\n",
"a(dynamic) = 3.140000\n",
"b(dynamic) = 5\n"
]
}
],
"source": [
"@cute.jit\n",
"def bar():\n",
" a = cutlass.Float32(3.14)\n",
" print(\"a(static) =\", a) # prints `a(static) = ?`\n",
" cute.printf(\"a(dynamic) = {}\", a) # prints `a(dynamic) = 3.140000`\n",
"\n",
" b = cutlass.Int32(5)\n",
" print(\"b(static) =\", b) # prints `b(static) = 5`\n",
" cute.printf(\"b(dynamic) = {}\", b) # prints `b(dynamic) = 5`\n",
"\n",
"bar()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Type Conversion API\n",
"\n",
"CUTLASS numeric types provide type conversion through the `to()` method available on all Numeric types. This allows you to convert between different numeric data types at runtime.\n",
"\n",
"Syntax:\n",
"\n",
"```python\n",
"new_value = value.to(target_type)\n",
"```\n",
"\n",
"The `to()` method supports conversion between:\n",
"- Integer types (Int8, Int16, Int32, Int64, UInt8, UInt16, UInt32, UInt64)\n",
"- Floating point types (Float16, Float32, Float64, BFloat16)\n",
"- Mixed integer/floating point conversions\n",
"\n",
"Note that when converting from floating point to integer types, the decimal portion is truncated. When converting between types with different ranges, values may be clamped or lose precision if they exceed the target type's representable range."
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Int32(42) => Float32(42.000000)\n",
"Float32(3.140000) => Int32(3)\n",
"Int32(127) => Int8(127)\n",
"Int32(300) => Int8(44) (truncated due to range limitation)\n"
]
}
],
"source": [
"@cute.jit\n",
"def type_conversion():\n",
" # Convert from Int32 to Float32\n",
" x = cutlass.Int32(42)\n",
" y = x.to(cutlass.Float32)\n",
" cute.printf(\"Int32({}) => Float32({})\", x, y)\n",
"\n",
" # Convert from Float32 to Int32\n",
" a = cutlass.Float32(3.14)\n",
" b = a.to(cutlass.Int32)\n",
" cute.printf(\"Float32({}) => Int32({})\", a, b)\n",
"\n",
" # Convert from Int32 to Int8\n",
" c = cutlass.Int32(127)\n",
" d = c.to(cutlass.Int8)\n",
" cute.printf(\"Int32({}) => Int8({})\", c, d)\n",
"\n",
" # Convert from Int32 to Int8 with value exceeding Int8 range\n",
" e = cutlass.Int32(300)\n",
" f = e.to(cutlass.Int8)\n",
" cute.printf(\"Int32({}) => Int8({}) (truncated due to range limitation)\", e, f)\n",
"\n",
"type_conversion()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Operator Overloading\n",
"\n",
"CUTLASS numeric types support Python's built-in operators, allowing you to write natural mathematical expressions. The operators work with both CUTLASS numeric types and Python native numeric types.\n",
"\n",
"Supported operators include:\n",
"- Arithmetic: `+`, `-`, `*`, `/`, `//`, `%`, `**`\n",
"- Comparison: `<`, `<=`, `==`, `!=`, `>=`, `>`\n",
"- Bitwise: `&`, `|`, `^`, `<<`, `>>`\n",
"- Unary: `-` (negation), `~` (bitwise NOT)"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"a: Int32(10), b: Int32(3)\n",
"x: Float32(5.500000)\n",
"\n",
"a + b = 13\n",
"x * 2 = 11.000000\n",
"a + x = 15.500000 (Int32 + Float32 promotes to Float32)\n",
"a / b = 3.333333\n",
"x / 2.0 = 2.750000\n",
"a > b = 1\n",
"a & b = 2\n",
"-a = -10\n",
"~a = -11\n"
]
}
],
"source": [
"@cute.jit\n",
"def operator_demo():\n",
" # Arithmetic operators\n",
" a = cutlass.Int32(10)\n",
" b = cutlass.Int32(3)\n",
" cute.printf(\"a: Int32({}), b: Int32({})\", a, b)\n",
"\n",
" x = cutlass.Float32(5.5)\n",
" cute.printf(\"x: Float32({})\", x)\n",
"\n",
" cute.printf(\"\")\n",
"\n",
" sum_result = a + b\n",
" cute.printf(\"a + b = {}\", sum_result)\n",
"\n",
" y = x * 2 # Multiplying with Python native type\n",
" cute.printf(\"x * 2 = {}\", y)\n",
"\n",
" # Mixed type arithmetic (Int32 + Float32) that integer is converted into float32\n",
" mixed_result = a + x\n",
" cute.printf(\"a + x = {} (Int32 + Float32 promotes to Float32)\", mixed_result)\n",
"\n",
" # Division with Int32 (note: integer division)\n",
" div_result = a / b\n",
" cute.printf(\"a / b = {}\", div_result)\n",
"\n",
" # Float division\n",
" float_div = x / cutlass.Float32(2.0)\n",
" cute.printf(\"x / 2.0 = {}\", float_div)\n",
"\n",
" # Comparison operators\n",
" is_greater = a > b\n",
" cute.printf(\"a > b = {}\", is_greater)\n",
"\n",
" # Bitwise operators\n",
" bit_and = a & b\n",
" cute.printf(\"a & b = {}\", bit_and)\n",
"\n",
" neg_a = -a\n",
" cute.printf(\"-a = {}\", neg_a)\n",
"\n",
" not_a = ~a\n",
" cute.printf(\"~a = {}\", not_a)\n",
"\n",
"operator_demo()\n"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.5"
}
},
"nbformat": 4,
"nbformat_minor": 4
}

View File

@@ -0,0 +1,838 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {
"editable": true,
"slideshow": {
"slide_type": ""
},
"tags": []
},
"outputs": [],
"source": [
"import torch\n",
"from functools import partial\n",
"\n",
"import cutlass\n",
"import cutlass.cute as cute\n",
"from cutlass.cute.runtime import from_dlpack"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Tutorial: Elementwise Add Kernel in CuTe DSL\n",
"\n",
"This tutorial demonstrates how to implement a simple elementwise\n",
"addition kernel using the CuTe DSL (Domain Specific Language).\n",
"\n",
"\n",
"\n",
"Elementwise Addition\n",
"---------------------\n",
"\n",
"Elementwise addition is a fundamental operation in linear algebra.\n",
"Given two tensors of the same shape, the operation performs element-wise\n",
"addition to produce a result tensor of the same shape.\n",
"\n",
"For two 2D tensors :math:`A` and :math:`B` of shape :math:`(M, N)`,\n",
"the elementwise addition operation :math:`C = A + B` is defined as:\n",
"\n",
"$\n",
" C_{i,j} = A_{i,j} + B_{i,j}\n",
"$\n",
"\n",
"where:\n",
"\n",
"- $i \\in [0, M-1]$ represents the row index\n",
"- $j \\in [0, N-1]$ represents the column index\n",
"- $A_{i,j}$, $B_{i,j}$, and $C_{i,j}$ are the elements at position $(i,j)$ \n",
" in tensors $A$, $B$, and $C$ respectively\n",
"\n",
"This operation is performed independently for each element position,\n",
"making it highly parallelizable and well-suited for GPU implementation.\n",
"\n",
"Naive Elementwise Add Kernel\n",
"-----------------------------\n",
"\n",
"Let's start with a naive implementation that loads each element from\n",
"$A$ and $B$, adds them, and stores the result back to $C$."
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [],
"source": [
"@cute.kernel\n",
"def naive_elementwise_add_kernel(\n",
" gA: cute.Tensor,\n",
" gB: cute.Tensor,\n",
" gC: cute.Tensor,\n",
"):\n",
" tidx, _, _ = cute.arch.thread_idx()\n",
" bidx, _, _ = cute.arch.block_idx()\n",
" bdim, _, _ = cute.arch.block_dim()\n",
"\n",
" thread_idx = bidx * bdim + tidx\n",
"\n",
" # Map thread index to logical index of input tensor\n",
" m, n = gA.shape\n",
" ni = thread_idx % n\n",
" mi = thread_idx // n\n",
"\n",
" # Map logical index to physical address via tensor layout\n",
" a_val = gA[mi, ni]\n",
" b_val = gB[mi, ni]\n",
"\n",
" # Perform element-wise addition\n",
" gC[mi, ni] = a_val + b_val"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Structure of the Kernel\n",
"\n",
"The naive kernel simply maps each thread to one element with a 1-to-1 mapping.\n",
"In this kernel, we don't use CuTe layout algebra but only use basic\n",
"addressing to index the tensor.\n",
"\n",
"We can launch the kernel with the following JIT function:"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [],
"source": [
"@cute.jit\n",
"def naive_elementwise_add(\n",
" mA: cute.Tensor,\n",
" mB: cute.Tensor,\n",
" mC: cute.Tensor\n",
"):\n",
" num_threads_per_block = 256\n",
"\n",
" m, n = mA.shape\n",
" kernel = naive_elementwise_add_kernel(mA, mB, mC)\n",
" kernel.launch(grid=((m * n) // num_threads_per_block, 1, 1),\n",
" block=(num_threads_per_block, 1, 1))\n",
"\n",
"M, N = 2048, 2048\n",
"\n",
"a = torch.randn(M, N, device=\"cuda\", dtype=torch.float16)\n",
"b = torch.randn(M, N, device=\"cuda\", dtype=torch.float16)\n",
"c = torch.zeros(M, N, device=\"cuda\", dtype=torch.float16)\n",
"\n",
"a_ = from_dlpack(a, assumed_align=16)\n",
"b_ = from_dlpack(b, assumed_align=16)\n",
"c_ = from_dlpack(c, assumed_align=16)\n",
"\n",
"# Compile kernel\n",
"naive_elementwise_add_ = cute.compile(naive_elementwise_add, a_, b_, c_)\n",
"naive_elementwise_add_(a_, b_, c_)\n",
"\n",
"# verify correctness\n",
"torch.testing.assert_close(c, a + b)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Benchmark performance\n",
"\n",
"Here's a utility function to benchmark our kernel implementations:"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [],
"source": [
"def benchmark(callable, *, num_warmups, num_iterations):\n",
" start_event = torch.cuda.Event(enable_timing=True)\n",
" end_event = torch.cuda.Event(enable_timing=True)\n",
"\n",
" torch.cuda.synchronize()\n",
"\n",
" for _ in range(num_warmups):\n",
" callable()\n",
"\n",
" start_event.record(stream=torch.cuda.current_stream())\n",
" for _ in range(num_iterations):\n",
" callable()\n",
" end_event.record(stream=torch.cuda.current_stream())\n",
" torch.cuda.synchronize()\n",
"\n",
" elapsed_time = start_event.elapsed_time(end_event)\n",
" avg_time = elapsed_time / num_iterations\n",
"\n",
" print(f\"Average execution time: {avg_time:.4f} ms\")\n",
" print(f\"Throughput: {(3 * a.numel() * 2) / (avg_time / 1000) / 1e9:.2f} GB/s\")"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Average execution time: 0.0385 ms\n",
"Throughput: 653.44 GB/s\n"
]
}
],
"source": [
"benchmark(partial(naive_elementwise_add_, a_, b_, c_), num_warmups=5, num_iterations=100)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Performance Analysis\n",
"\n",
"While our naive implementation maps thread indices to contiguous tensor\n",
"dimensions for coalesced memory access, it doesn't have enough\n",
"in-flight load & store operations to hide memory latency.\n",
"\n",
"According to Little's Law:\n",
"\n",
"$ L = \\lambda \\times W $\n",
"\n",
"Where:\n",
"- $L$ is the average number of items in a system\n",
"- $\\lambda$ is the average arrival rate of items (bandwidth)\n",
"- $W$ is the average time an item spends in the system (latency)\n",
"\n",
"For our elementwise addition kernel:\n",
"\n",
"1. $L$: The number of load & store operations in-flight\n",
"2. $\\lambda$ (Bandwidth): Data transfer rate between memory and compute units\n",
"3. $W$ (Latency): Round-trip delay of memory requests\n",
"\n",
"For memory-bound operations like elementwise addition, performance is\n",
"limited by the number of in-flight load & store operations.\n",
"\n",
"## Vectorized Load and Store\n",
"\n",
"To improve performance according to Little's Law, we need to increase the number\n",
"of in-flight requests. We can do this by increasing the number of bytes handled\n",
"in each load & store operation per thread through vectorized memory access.\n",
"\n",
"Since Ampere GPUs support up to 128-bit per load/store and each element is 32-bit,\n",
"we can load 4 elements per vectorized operation on contiguous rows.\n",
"CuTe tiling operations make this vectorization straightforward.\n",
"\n",
"Using ``tiled_tensor = cute.zipped_divide(tensor, tiler)``, we can partition the input\n",
"``tensor`` into groups of ``tiler`` blocks. For vectorization, we specify ``tiler``\n",
"as the block of data each thread accesses (4 contiguous elements in the same row, or ``(1,4)``).\n",
"Different threads can then access different blocks by indexing into the 2nd mode of ``tiled_tensor``.\n",
"\n",
"```python\n",
"mA : cute.Tensor # (2048,2048):(2048,1)\n",
"gA = cute.zipped_divide(a, tiler=(1, 4)) # tiled/vectorized => ((1,4),(2048,512)):((0,1),(2048,4))\n",
"```\n",
"\n",
"$\n",
" \\begin{array}{ccccc}\n",
" & ((1,4) & , & (2048,512)) & : ((0,1),(2048,4)) \\\\\n",
" & \\underbrace{\\phantom{(1,4)}}_{tiler} & & \\underbrace{\\phantom{(2048,512)}}_{threads} & \\\\\n",
" & \\text{\\scriptsize per-thread} & & \\text{\\scriptsize num of tiles}\n",
" \\end{array}\n",
"$"
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [],
"source": [
"@cute.kernel\n",
"def vectorized_elementwise_add_kernel(\n",
" gA: cute.Tensor,\n",
" gB: cute.Tensor,\n",
" gC: cute.Tensor,\n",
"):\n",
" tidx, _, _ = cute.arch.thread_idx()\n",
" bidx, _, _ = cute.arch.block_idx()\n",
" bdim, _, _ = cute.arch.block_dim()\n",
"\n",
" thread_idx = bidx * bdim + tidx\n",
"\n",
" # Map thread index to logical index of input tensor\n",
" m, n = gA.shape[1] # thread-domain\n",
" ni = thread_idx % n\n",
" mi = thread_idx // n\n",
"\n",
" # Map logical index to physical address via tensor layout\n",
" a_val = gA[(None, (mi, ni))].load()\n",
" b_val = gB[(None, (mi, ni))].load()\n",
" print(f\"[DSL INFO] sliced gA = {gA[(None, (mi, ni))]}\")\n",
" print(f\"[DSL INFO] sliced gB = {gB[(None, (mi, ni))]}\")\n",
"\n",
" # Perform element-wise addition\n",
" gC[(None, (mi, ni))] = a_val + b_val"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"This vectorized kernel follows a similar structure to its naive non-vectorized counterpart,\n",
"with one key difference: the tensor slicing pattern. By using `(None, (mi, ni))` as the slice indices,\n",
"we can extract a `(1,4)` sub-tensor from `gA`, `gB` and `gC` like \n",
"\n",
"```python\n",
"gA[(None, (mi, ni))]\n",
"\n",
"```\n",
"\n",
"Then tensor data can be loaded into vector via the `.load()` method.\n",
"\n",
"\n",
"```\n",
" slice\n",
" ((1,4),(2048,512)):((0,1),(2048,4)) ==> ((1,4)):((0,1))\n",
" ^ ^ ^\n",
" | | |\n",
" (None, (mi, ni))\n",
"```"
]
},
{
"cell_type": "code",
"execution_count": 7,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"[DSL INFO] Tiled Tensors:\n",
"[DSL INFO] gA = tensor<ptr<f16, gmem, align<16>> o ((1,4),(2048,512)):((0,1),(2048,4))>\n",
"[DSL INFO] gB = tensor<ptr<f16, gmem, align<16>> o ((1,4),(2048,512)):((0,1),(2048,4))>\n",
"[DSL INFO] gC = tensor<ptr<f16, gmem, align<16>> o ((1,4),(2048,512)):((0,1),(2048,4))>\n",
"[DSL INFO] sliced gA = tensor<ptr<f16, gmem, align<8>> o ((1,4)):((0,1))>\n",
"[DSL INFO] sliced gB = tensor<ptr<f16, gmem, align<8>> o ((1,4)):((0,1))>\n"
]
}
],
"source": [
"@cute.jit\n",
"def vectorized_elementwise_add(\n",
" mA: cute.Tensor,\n",
" mB: cute.Tensor,\n",
" mC: cute.Tensor\n",
"):\n",
" threads_per_block = 256\n",
"\n",
" gA = cute.zipped_divide(mA, (1, 4))\n",
" gB = cute.zipped_divide(mB, (1, 4))\n",
" gC = cute.zipped_divide(mC, (1, 4))\n",
"\n",
" print(f\"[DSL INFO] Tiled Tensors:\")\n",
" print(f\"[DSL INFO] gA = {gA}\")\n",
" print(f\"[DSL INFO] gB = {gB}\")\n",
" print(f\"[DSL INFO] gC = {gC}\")\n",
"\n",
" vectorized_elementwise_add_kernel(gA, gB, gC).launch(\n",
" grid=(cute.size(gC, mode=[1]) // threads_per_block, 1, 1),\n",
" block=(threads_per_block, 1, 1),\n",
" )\n",
"\n",
"a = torch.randn(M, N, device=\"cuda\", dtype=torch.float16)\n",
"b = torch.randn(M, N, device=\"cuda\", dtype=torch.float16)\n",
"c = torch.zeros(M, N, device=\"cuda\", dtype=torch.float16)\n",
"\n",
"a_ = from_dlpack(a, assumed_align=16)\n",
"b_ = from_dlpack(b, assumed_align=16)\n",
"c_ = from_dlpack(c, assumed_align=16)\n",
"\n",
"compiled_func = cute.compile(vectorized_elementwise_add, a_, b_, c_)\n",
"compiled_func(a_, b_, c_)\n",
"\n",
"# verify correctness\n",
"torch.testing.assert_close(c, a + b)"
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Average execution time: 0.0202 ms\n",
"Throughput: 1244.98 GB/s\n"
]
}
],
"source": [
"benchmark(partial(compiled_func, a_, b_, c_), num_warmups=5, num_iterations=100)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## TV Layout\n",
"\n",
"Both the naive and vectorized kernels follow a common pattern to map thread indices\n",
"to physical addresses:\n",
"\n",
"Step 1: Map thread index to logical M/N coordinates\n",
"\n",
"```python\n",
" mi = thread_idx // n\n",
" ni = thread_idx % n\n",
"```\n",
"\n",
"Step 2: Map logical M/N coordinates to physical addresses using the tensor layout\n",
"\n",
"```python\n",
" a[(None, (mi, ni))].load()\n",
"```\n",
"\n",
"CuTe uses TV layout to represent this mapping from thread index and value index\n",
"(i.e., the 4 elements loaded per thread) to the logical coordinate space of a tensor.\n",
"By configuring different TV layouts, we can experiment with different memory access\n",
"patterns with minimal code changes.\n",
"\n",
"The following example demonstrates two levels of tiling: at the thread-block level\n",
"and at the thread level.\n",
"\n",
"For thread-block level tiling, each input & output tensor is first divided\n",
"into a group of ``(TileM, TileN)`` sub-tensors at the host side.\n",
"\n",
"Inside the GPU kernel, we provide the thread-block index to the 2nd mode of the tiled tensor\n",
"(``gA[((None, None), bidx)]``), which returns a thread-block local view of\n",
"a single ``(TileM, TileN)`` sub-tensor.\n",
"\n",
"For thread level tiling, we compose the sub-tensor (which maps from logical coordinates\n",
"to physical addresses) with the TV layout (which maps from thread & value indices to\n",
"logical coordinates). This gives us a tiled sub-tensor that maps from thread & value\n",
"indices directly to physical addresses.\n",
"\n",
"We then provide the thread index to the tiled sub-tensor (``tidfrgA[(tidx, None)]``)\n",
"to get a thread-local view of the data each thread accesses. Note that the thread index\n",
"is now in the 1st mode, as the tiled sub-tensor puts the thread mode before the value mode."
]
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {},
"outputs": [],
"source": [
"@cute.kernel\n",
"def elementwise_add_kernel(\n",
" gA: cute.Tensor,\n",
" gB: cute.Tensor,\n",
" gC: cute.Tensor,\n",
" tv_layout: cute.Layout\n",
"):\n",
" tidx, _, _ = cute.arch.thread_idx()\n",
" bidx, _, _ = cute.arch.block_idx()\n",
"\n",
" #--------------------------------\n",
" # slice for thread-block level view\n",
" #--------------------------------\n",
" blk_coord = ((None, None), bidx)\n",
"\n",
" # logical coord -> address\n",
" blkA = gA[blk_coord] # (TileM, TileN) -> physical address\n",
" blkB = gB[blk_coord] # (TileM, TileN) -> physical address\n",
" blkC = gC[blk_coord] # (TileM, TileN) -> physical address\n",
"\n",
" #--------------------------------\n",
" # compose for thread-index & value-index to physical mapping\n",
" #--------------------------------\n",
" # blockA: (TileM, TileN) -> physical address\n",
" # tv_layout: (tid, vid) -> (TileM, TileN)\n",
" # tidfrgA = blkA o tv_layout\n",
" # tidfrgA: (tid, vid) -> physical address\n",
" tidfrgA = cute.composition(blkA, tv_layout)\n",
" tidfrgB = cute.composition(blkB, tv_layout)\n",
" tidfrgC = cute.composition(blkC, tv_layout)\n",
"\n",
" print(f\"Composed with TV layout:\")\n",
" print(f\" tidfrgA: {tidfrgA.type}\")\n",
"\n",
" #--------------------------------\n",
" # slice for thread-level view\n",
" #--------------------------------\n",
" # `None` represent slice of the entire per-thread data\n",
" thr_coord = (tidx, None)\n",
"\n",
" # slice for threads: vid -> address\n",
" thrA = tidfrgA[thr_coord] # (V) -> physical address\n",
" thrB = tidfrgB[thr_coord] # (V) -> physical address\n",
" thrC = tidfrgC[thr_coord] # (V) -> physical address\n",
"\n",
" thrC[None] = thrA.load() + thrB.load()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"If we take a closer look at the layout of zipped divided input tensor `gA`:\n",
"\n",
"```\n",
"Tiled to Thread Block:\n",
"\n",
" ((16,256),(128,8)) : ((2048,1),(32768,256))\n",
" ~~~~~~~~ ~~~~~~ ~~~~~~~~\n",
" | | |\n",
" | | |\n",
" | `------------------------> Number of Thread Blocks\n",
" | |\n",
" | |\n",
" `--------------------'\n",
" |\n",
" V\n",
" Thread Block\n",
" Tile\n",
"\n",
"Sliced to Thread-Block local sub-tensor (a (16, 128) tile): gA[((None, None), bidx)]\n",
"\n",
" (16,256) : (2048,1)\n",
" ~~~~~~ ~~~~~~\n",
" | | Tiled/Composed with TV Layout\n",
" | | \n",
" | | o ((32,4),(8,4)):((128,4),(16,1))\n",
" V V \n",
"~~~~~~~~~~~~~~~ ~~~~~~~~~~~~~~~~~~~ \n",
"((32,4), (8,4)) : ((4,8192),(1,2048))\n",
" | |\n",
" | `--------> per thread fragment\n",
" |\n",
"Thread Block\n",
" Shape\n",
"\n",
"Sliced to Thread local sub-tensor (a (4,8) tile): tidfrgA[(tidx, None)]\n",
"\n",
"```\n",
"\n",
"The host code below shows the construction of the TV layout. By composing\n",
"a thread layout of ``(4,32):(32,1)`` (32 threads read contiguous elements on the row dimension,\n",
"then 4 warps read different rows) with a value layout of ``(4,8):(8,1)`` (each thread reads\n",
"8 contiguous elements on the row dimension across 4 contiguous rows),\n",
"we obtain the TV layout shown in the figure above."
]
},
{
"cell_type": "code",
"execution_count": 10,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Tiler: (16, 256)\n",
"TV Layout: ((32,4),(8,4)):((128,4),(16,1))\n",
"Tiled Input Tensors:\n",
" gA: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
" gB: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
" gC: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
"Composed with TV layout:\n",
" tidfrgA: !cute.memref<f16, gmem, align<16>, \"((32,4),(8,4)):((8,8192),(1,2048))\">\n"
]
}
],
"source": [
"@cute.jit\n",
"def elementwise_add(\n",
" mA: cute.Tensor,\n",
" mB: cute.Tensor,\n",
" mC: cute.Tensor,\n",
"):\n",
" # mA layout: (M, N):(N, 1)\n",
" # TV layout map thread & value index to (16, 256) logical tile\n",
" # - contiguous thread index maps to mode-1 because input layout is contiguous on\n",
" # mode-1 for coalesced load-store\n",
" # - each thread load 8 contiguous element each row and load 4 rows\n",
" thr_layout = cute.make_layout((4, 32), stride=(32, 1))\n",
" val_layout = cute.make_layout((4, 8), stride=(8, 1))\n",
" tiler_mn, tv_layout = cute.make_layout_tv(thr_layout, val_layout)\n",
" print(f\"Tiler: {tiler_mn}\")\n",
" print(f\"TV Layout: {tv_layout}\")\n",
"\n",
" gA = cute.zipped_divide(mA, tiler_mn) # ((TileM, TileN), (RestM, RestN))\n",
" gB = cute.zipped_divide(mB, tiler_mn) # ((TileM, TileN), (RestM, RestN))\n",
" gC = cute.zipped_divide(mC, tiler_mn) # ((TileM, TileN), (RestM, RestN))\n",
"\n",
" print(f\"Tiled Input Tensors:\")\n",
" print(f\" gA: {gA.type}\")\n",
" print(f\" gB: {gB.type}\")\n",
" print(f\" gC: {gC.type}\")\n",
"\n",
" # Launch the kernel asynchronously\n",
" # Async token(s) can also be specified as dependencies\n",
" elementwise_add_kernel(\n",
" gA, gB, gC, tv_layout\n",
" ).launch(\n",
" grid=[cute.size(gC, mode=[1]), 1, 1],\n",
" block=[cute.size(tv_layout, mode=[0]), 1, 1],\n",
" )\n",
"\n",
"a = torch.randn(M, N, device=\"cuda\", dtype=torch.float16)\n",
"b = torch.randn(M, N, device=\"cuda\", dtype=torch.float16)\n",
"c = torch.zeros(M, N, device=\"cuda\", dtype=torch.float16)\n",
"\n",
"a_ = from_dlpack(a, assumed_align=16)\n",
"b_ = from_dlpack(b, assumed_align=16)\n",
"c_ = from_dlpack(c, assumed_align=16)\n",
"\n",
"elementwise_add_ = cute.compile(elementwise_add, a_, b_, c_)\n",
"elementwise_add_(a_, b_, c_)\n",
"\n",
"# verify correctness\n",
"torch.testing.assert_close(c, a + b)"
]
},
{
"cell_type": "code",
"execution_count": 11,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Average execution time: 0.0222 ms\n",
"Throughput: 1133.58 GB/s\n"
]
}
],
"source": [
"benchmark(partial(elementwise_add_, a_, b_, c_), num_warmups=5, num_iterations=200)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Using Lambda Function\n",
"\n",
"CuTe DSL is built on top of Python. It can leverage Python to implement meta-programming to generate flexible kernels.\n",
"E.g. we can write kernel template that take custom binary operations to generate kernels for arbitrary binary operations.\n",
"\n",
"\n",
"```python\n",
"@cute.jit\n",
"def elementwise_apply(\n",
" op: cutlass.Constexpr,\n",
" mA: cute.Tensor,\n",
" mB: cute.Tensor,\n",
" mC: cute.Tensor\n",
"):\n",
" ...\n",
"\n",
"```"
]
},
{
"cell_type": "code",
"execution_count": 12,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Tiler: (16, 256)\n",
"TV Layout: ((32,4),(8,4)):((128,4),(16,1))\n",
"Tiled Input Tensors:\n",
" gA: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
" gB: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
" gC: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
"Composed with TV layout:\n",
" tidfrgA: !cute.memref<f16, gmem, align<16>, \"((32,4),(8,4)):((8,8192),(1,2048))\">\n"
]
}
],
"source": [
"@cute.kernel\n",
"def elementwise_apply_kernel(\n",
" op: cutlass.Constexpr, # lambda function must be const expr to generate code at compile time\n",
" gA: cute.Tensor,\n",
" gB: cute.Tensor,\n",
" gC: cute.Tensor,\n",
" tv_layout: cute.Layout\n",
"):\n",
" tidx, _, _ = cute.arch.thread_idx()\n",
" bidx, _, _ = cute.arch.block_idx()\n",
"\n",
" blk_coord = ((None, None), bidx)\n",
"\n",
" # logical coord -> address\n",
" blkA = gA[blk_coord] # (TileM, TileN) -> physical address\n",
" blkB = gB[blk_coord] # (TileM, TileN) -> physical address\n",
" blkC = gC[blk_coord] # (TileM, TileN) -> physical address\n",
"\n",
" tidfrgA = cute.composition(blkA, tv_layout)\n",
" tidfrgB = cute.composition(blkB, tv_layout)\n",
" tidfrgC = cute.composition(blkC, tv_layout)\n",
"\n",
" print(f\"Composed with TV layout:\")\n",
" print(f\" tidfrgA: {tidfrgA.type}\")\n",
"\n",
" thr_coord = (tidx, None)\n",
"\n",
" # slice for threads: vid -> address\n",
" thrA = tidfrgA[thr_coord] # (V) -> physical address\n",
" thrB = tidfrgB[thr_coord] # (V) -> physical address\n",
" thrC = tidfrgC[thr_coord] # (V) -> physical address\n",
"\n",
" #--------------------------------\n",
" # apply custom operation\n",
" #--------------------------------\n",
" thrC[None] = op(thrA.load(), thrB.load())\n",
"\n",
"\n",
"@cute.jit\n",
"def elementwise_op(\n",
" op: cutlass.Constexpr,\n",
" mA: cute.Tensor,\n",
" mB: cute.Tensor,\n",
" mC: cute.Tensor,\n",
"):\n",
" # mA layout: (M, N):(N, 1)\n",
" # TV layout map thread & value index to (16, 256) logical tile\n",
" # - contiguous thread index maps to mode-1 because input layout is contiguous on\n",
" # mode-1 for coalesced load-store\n",
" # - each thread load 8 contiguous element each row and load 4 rows\n",
" thr_layout = cute.make_layout((4, 32), stride=(32, 1))\n",
" val_layout = cute.make_layout((4, 8), stride=(8, 1))\n",
" tiler_mn, tv_layout = cute.make_layout_tv(thr_layout, val_layout)\n",
" print(f\"Tiler: {tiler_mn}\")\n",
" print(f\"TV Layout: {tv_layout}\")\n",
"\n",
" gA = cute.zipped_divide(mA, tiler_mn) # ((TileM, TileN), (RestM, RestN))\n",
" gB = cute.zipped_divide(mB, tiler_mn) # ((TileM, TileN), (RestM, RestN))\n",
" gC = cute.zipped_divide(mC, tiler_mn) # ((TileM, TileN), (RestM, RestN))\n",
"\n",
" print(f\"Tiled Input Tensors:\")\n",
" print(f\" gA: {gA.type}\")\n",
" print(f\" gB: {gB.type}\")\n",
" print(f\" gC: {gC.type}\")\n",
"\n",
" # Launch the kernel asynchronously\n",
" # Async token(s) can also be specified as dependencies\n",
" elementwise_apply_kernel(\n",
" op, gA, gB, gC, tv_layout\n",
" ).launch(\n",
" grid=[cute.size(gC, mode=[1]), 1, 1],\n",
" block=[cute.size(tv_layout, mode=[0]), 1, 1],\n",
" )\n",
"\n",
"a = torch.randn(M, N, device=\"cuda\", dtype=torch.float16)\n",
"b = torch.randn(M, N, device=\"cuda\", dtype=torch.float16)\n",
"c = torch.zeros(M, N, device=\"cuda\", dtype=torch.float16)\n",
"\n",
"a_ = from_dlpack(a, assumed_align=16)\n",
"b_ = from_dlpack(b, assumed_align=16)\n",
"c_ = from_dlpack(c, assumed_align=16)\n",
"\n",
"from operator import mul\n",
"\n",
"elementwise_op(mul, a_, b_, c_)\n",
"\n",
"# verify correctness\n",
"torch.testing.assert_close(c, mul(a, b))"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"Custom operators can be more complex. For example, here's a function that performs\n",
"multiplication followed by ReLU:"
]
},
{
"cell_type": "code",
"execution_count": 13,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Tiler: (16, 256)\n",
"TV Layout: ((32,4),(8,4)):((128,4),(16,1))\n",
"Tiled Input Tensors:\n",
" gA: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
" gB: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
" gC: !cute.memref<f16, gmem, align<16>, \"((16,256),(128,8)):((2048,1),(32768,256))\">\n",
"Composed with TV layout:\n",
" tidfrgA: !cute.memref<f16, gmem, align<16>, \"((32,4),(8,4)):((8,8192),(1,2048))\">\n"
]
}
],
"source": [
"def mul_relu(a, b):\n",
" tmp = a * b\n",
" return cute.where(tmp > 0, tmp, cute.full_like(tmp, 0))\n",
"\n",
"\n",
"# As we uses cute.where in customized operation, we need to create another relu function\n",
"def mul_relu_ref(a, b):\n",
" tmp = a * b\n",
" return torch.relu(tmp)\n",
"\n",
"\n",
"elementwise_op(mul_relu, a_, b_, c_)\n",
"\n",
"# verify correctness\n",
"torch.testing.assert_close(c, mul_relu_ref(a, b))"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.5"
},
"widgets": {
"application/vnd.jupyter.widget-state+json": {
"state": {},
"version_major": 2,
"version_minor": 0
}
}
},
"nbformat": 4,
"nbformat_minor": 4
}

View File

@@ -0,0 +1,173 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Your First Program with CuTe DSL\n",
"\n",
"## Introduction\n",
"\n",
"Welcome! In this tutorial, we'll write a simple \"Hello World\" program that runs on your GPU using CuTe DSL. This will help you understand the basics of GPU programming with our framework.\n",
"\n",
"### What You'll Learn\n",
"\n",
"- How to write code that runs on both CPU (host) and GPU (device),\n",
"- How to launch a GPU kernel (a function that runs on the GPU),\n",
"- Basic CUDA concepts like threads and thread blocks,\n",
"\n",
"### Step 1: Import Required Libraries\n",
"\n",
"First, let's import the libraries we need:"
]
},
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [],
"source": [
"import cutlass \n",
"import cutlass.cute as cute "
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"\n",
"### Step 2: Write Our GPU Kernel\n",
"A GPU kernel is a function that runs on the GPU. Here's a simple kernel that prints \"Hello World\".\n",
"Key concepts:\n",
"- `@cute.kernel`: This decorator tells CUTLASS that this function should run on the GPU\n",
"- `cute.arch.thread_idx()`: Gets the ID of the current GPU thread (like a worker's ID number)\n",
"- We only want one thread to print the message (thread 0) to avoid multiple prints"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [],
"source": [
"@cute.kernel\n",
"def kernel():\n",
" # Get the x component of the thread index (y and z components are unused)\n",
" tidx, _, _ = cute.arch.thread_idx()\n",
" # Only the first thread (thread 0) prints the message\n",
" if tidx == 0:\n",
" cute.printf(\"Hello world\")"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Step 3: Write Our Host Function\n",
"\n",
"Now we need a function that sets up the GPU and launches our kernel.\n",
"Key concepts:\n",
"- `@cute.jit`: This decorator is for functions that run on the CPU but can launch GPU code\n",
"- We need to initialize CUDA before using the GPU\n",
"- `.launch()` tells CUDA how many blocks, threads, shared memory, etc. to use"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [],
"source": [
"@cute.jit\n",
"def hello_world():\n",
"\n",
" # Print hello world from host code\n",
" cute.printf(\"hello world\")\n",
" \n",
" # Initialize CUDA context for launching a kernel with error checking\n",
" # We make context initialization explicit to allow users to control the context creation \n",
" # and avoid potential issues with multiple contexts\n",
" cutlass.cuda.initialize_cuda_context()\n",
"\n",
" # Launch kernel\n",
" kernel().launch(\n",
" grid=(1, 1, 1), # Single thread block\n",
" block=(32, 1, 1) # One warp (32 threads) per thread block\n",
" )"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Step 4: Run Our Program\n",
"\n",
"There are 2 ways we can run our program:\n",
"\n",
"1. compile and run immediately\n",
"2. separate compilation which allows us to compile the code once and run multiple times\n",
" \n",
"Please note the `Compiling...` for Method 2 prints before the \"Hello world\" of the first kernel. This shows the asynchronous behavior between CPU and GPU prints. "
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Running hello_world()...\n",
"hello world\n",
"Compiling...\n",
"Hello world\n",
"Running compiled version...\n",
"hello world\n"
]
}
],
"source": [
"# Method 1: Just-In-Time (JIT) compilation - compiles and runs the code immediately\n",
"print(\"Running hello_world()...\")\n",
"hello_world()\n",
"\n",
"# Method 2: Compile first (useful if you want to run the same code multiple times)\n",
"print(\"Compiling...\")\n",
"hello_world_compiled = cute.compile(hello_world)\n",
"# Run the pre-compiled version\n",
"print(\"Running compiled version...\")\n",
"hello_world_compiled()"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.5"
},
"widgets": {
"application/vnd.jupyter.widget-state+json": {
"state": {},
"version_major": 2,
"version_minor": 0
}
}
},
"nbformat": 4,
"nbformat_minor": 4
}

Binary file not shown.

After

Width:  |  Height:  |  Size: 8.4 KiB

View File

@@ -0,0 +1,425 @@
{
"cells": [
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Printing with CuTe DSL\n",
"\n",
"This notebook demonstrates the different ways to print values in CuTe and explains the important distinction between static (compile-time) and dynamic (runtime) values.\n",
"\n",
"## Key Concepts\n",
"- Static values: Known at compile time\n",
"- Dynamic values: Only known at runtime\n",
"- Different printing methods for different scenarios\n",
"- Layout representation in CuTe\n",
"- Tensor visualization and formatting"
]
},
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [],
"source": [
"import cutlass\n",
"import cutlass.cute as cute\n",
"import numpy as np"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Print Example Function\n",
"\n",
"The `print_example` function demonstrates several important concepts:\n",
"\n",
"### 1. Python's `print` vs CuTe's `cute.printf`\n",
"- `print`: Can only show static values at compile time\n",
"- `cute.printf`: Can display both static and dynamic values at runtime\n",
"\n",
"### 2. Value Types\n",
"- `a`: Dynamic `Int32` value (runtime)\n",
"- `b`: Static `Constexpr[int]` value (compile-time)\n",
"\n",
"### 3. Layout Printing\n",
"Shows how layouts are represented differently in static vs dynamic contexts:\n",
"- Static context: Unknown values shown as `?`\n",
"- Dynamic context: Actual values displayed"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [],
"source": [
"@cute.jit\n",
"def print_example(a: cutlass.Int32, b: cutlass.Constexpr[int]):\n",
" \"\"\"\n",
" Demonstrates different printing methods in CuTe and how they handle static vs dynamic values.\n",
"\n",
" This example shows:\n",
" 1. How Python's `print` function works with static values at compile time but can't show dynamic values\n",
" 2. How `cute.printf` can display both static and dynamic values at runtime\n",
" 3. The difference between types in static vs dynamic contexts\n",
" 4. How layouts are represented in both printing methods\n",
"\n",
" Args:\n",
" a: A dynamic Int32 value that will be determined at runtime\n",
" b: A static (compile-time constant) integer value\n",
" \"\"\"\n",
" # Use Python `print` to print static information\n",
" print(\">>>\", b) # => 2\n",
" # `a` is dynamic value\n",
" print(\">>>\", a) # => ?\n",
"\n",
" # Use `cute.printf` to print dynamic information\n",
" cute.printf(\">?? {}\", a) # => 8\n",
" cute.printf(\">?? {}\", b) # => 2\n",
"\n",
" print(\">>>\", type(a)) # => <class 'cutlass.Int32'>\n",
" print(\">>>\", type(b)) # => <class 'int'>\n",
"\n",
" layout = cute.make_layout((a, b))\n",
" print(\">>>\", layout) # => (?,2):(1,?)\n",
" cute.printf(\">?? {}\", layout) # => (8,2):(1,8)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Compile and Run\n",
"\n",
"**Direct Compilation and Run**\n",
" - `print_example(cutlass.Int32(8), 2)`\n",
" - Compiles and runs in one step will execute both static and dynamic print\n",
" * `>>>` stands for static print\n",
" * `>??` stands for dynamic print"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
">>> 2\n",
">>> ?\n",
">>> Int32\n",
">>> <class 'int'>\n",
">>> (?,2):(1,?)\n",
">?? 8\n",
">?? 2\n",
">?? (8,2):(1,8)\n"
]
}
],
"source": [
"print_example(cutlass.Int32(8), 2)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Compile Function\n",
"\n",
"When compiles the function with `cute.compile(print_example, cutlass.Int32(8), 2)`, Python interpreter \n",
"traces code and only evaluate static expression and print static information."
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
">>> 2\n",
">>> ?\n",
">>> Int32\n",
">>> <class 'int'>\n",
">>> (?,2):(1,?)\n"
]
}
],
"source": [
"print_example_compiled = cute.compile(print_example, cutlass.Int32(8), 2)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Call compiled function\n",
"\n",
"Only print out runtime information"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
">?? 8\n",
">?? 2\n",
">?? (8,2):(1,8)\n"
]
}
],
"source": [
"print_example_compiled(cutlass.Int32(8))"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Format String Example\n",
"\n",
"The `format_string_example` function shows an important limitation:\n",
"- F-strings in CuTe are evaluated at compile time\n",
"- This means dynamic values won't show their runtime values in f-strings\n",
"- Use `cute.printf` when you need to see runtime values"
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Direct run output:\n",
"a: ?, b: 2\n",
"layout: (?,2):(1,?)\n"
]
}
],
"source": [
"@cute.jit\n",
"def format_string_example(a: cutlass.Int32, b: cutlass.Constexpr[int]):\n",
" \"\"\"\n",
" Format string is evaluated at compile time.\n",
" \"\"\"\n",
" print(f\"a: {a}, b: {b}\")\n",
"\n",
" layout = cute.make_layout((a, b))\n",
" print(f\"layout: {layout}\")\n",
"\n",
"print(\"Direct run output:\")\n",
"format_string_example(cutlass.Int32(8), 2)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Printing Tensor Examples\n",
"\n",
"CuTe provides specialized functionality for printing tensors through the `print_tensor` operation. The `cute.print_tensor` takes the following parameter:\n",
"- `Tensor` (required): A CuTe tensor object that you want to print. The tensor must support load and store operations\n",
"- `verbose` (optional, default=False): A boolean flag that controls the level of detail in the output. When set to True, it will print indices details for each element in the tensor.\n",
"\n",
"Below example code shows the difference between verbose ON and OFF, and how to print a sub range of the given tensor."
]
},
{
"cell_type": "code",
"execution_count": 7,
"metadata": {},
"outputs": [],
"source": [
"from cutlass.cute.runtime import from_dlpack\n",
"\n",
"@cute.jit\n",
"def print_tensor_basic(x : cute.Tensor):\n",
" # Print the tensor\n",
" print(\"Basic output:\")\n",
" cute.print_tensor(x)\n",
" \n",
"@cute.jit\n",
"def print_tensor_verbose(x : cute.Tensor):\n",
" # Print the tensor with verbose mode\n",
" print(\"Verbose output:\")\n",
" cute.print_tensor(x, verbose=True)\n",
"\n",
"@cute.jit\n",
"def print_tensor_slice(x : cute.Tensor, coord : tuple):\n",
" # slice a 2D tensor from the 3D tensor\n",
" sliced_data = cute.slice_(x, coord)\n",
" y = cute.make_fragment(sliced_data.layout, sliced_data.element_type)\n",
" # Convert to TensorSSA format by loading the sliced data into the fragment\n",
" y.store(sliced_data.load())\n",
" print(\"Slice output:\")\n",
" cute.print_tensor(y)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"The default `cute.print_tensor` will output CuTe tensor with datatype, storage space, CuTe layout information, and print data in torch-style format."
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Basic output:\n",
"tensor(raw_ptr(0x000000000a5f1d50: f32, generic, align<4>) o (4,3,2):(6,2,1), data=\n",
" [[[ 0.000000, 2.000000, 4.000000, ],\n",
" [ 6.000000, 8.000000, 10.000000, ],\n",
" [ 12.000000, 14.000000, 16.000000, ],\n",
" [ 18.000000, 20.000000, 22.000000, ]],\n",
"\n",
" [[ 1.000000, 3.000000, 5.000000, ],\n",
" [ 7.000000, 9.000000, 11.000000, ],\n",
" [ 13.000000, 15.000000, 17.000000, ],\n",
" [ 19.000000, 21.000000, 23.000000, ]]])\n"
]
}
],
"source": [
"def tensor_print_example1():\n",
" shape = (4, 3, 2)\n",
" \n",
" # Creates [0,...,23] and reshape to (4, 3, 2)\n",
" data = np.arange(24, dtype=np.float32).reshape(*shape) \n",
" \n",
" print_tensor_basic(from_dlpack(data))\n",
"\n",
"tensor_print_example1()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"The verbosed print will show coodination details of each element in the tensor. The below example shows how we index element in a 2D 4x3 tensor space."
]
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Verbose output:\n",
"tensor(raw_ptr(0x000000000a814cc0: f32, generic, align<4>) o (4,3):(3,1), data= (\n",
"\t(0,0)= 0.000000\n",
"\t(0,1)= 1.000000\n",
"\t(0,2)= 2.000000\n",
"\t(1,0)= 3.000000\n",
"\t(1,1)= 4.000000\n",
"\t(1,2)= 5.000000\n",
"\t(2,0)= 6.000000\n",
"\t(2,1)= 7.000000\n",
"\t(2,2)= 8.000000\n",
"\t(3,0)= 9.000000\n",
"\t(3,1)= 10.000000\n",
"\t(3,2)= 11.000000\n",
")\n"
]
}
],
"source": [
"def tensor_print_example2():\n",
" shape = (4, 3)\n",
" \n",
" # Creates [0,...,11] and reshape to (4, 3)\n",
" data = np.arange(12, dtype=np.float32).reshape(*shape) \n",
" \n",
" print_tensor_verbose(from_dlpack(data))\n",
"\n",
"tensor_print_example2()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"To print a subset elements in the given Tensor, we can use cute.slice_ to select a range of the given tensor, load them into register and then print the values with `cute.print_tensor`."
]
},
{
"cell_type": "code",
"execution_count": 10,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"Slice output:\n",
"tensor(raw_ptr(0x00007ffeeae1fc60: f32, rmem, align<32>) o (4):(3), data=\n",
" [ 0.000000, ],\n",
" [ 3.000000, ],\n",
" [Slice output:\n",
" 6.000000, ],\n",
" [ 9.000000, ])\n",
"tensor(raw_ptr(0x00007ffeeae1fc60: f32, rmem, align<32>) o (3):(1), data=\n",
" [ 3.000000, ],\n",
" [ 4.000000, ],\n",
" [ 5.000000, ])\n"
]
}
],
"source": [
"def tensor_print_example3():\n",
" shape = (4, 3)\n",
" \n",
" # Creates [0,...,11] and reshape to (4, 3)\n",
" data = np.arange(12, dtype=np.float32).reshape(*shape) \n",
" \n",
" print_tensor_slice(from_dlpack(data), (None, 0))\n",
" print_tensor_slice(from_dlpack(data), (1, None))\n",
"\n",
"tensor_print_example3()"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.5"
}
},
"nbformat": 4,
"nbformat_minor": 4
}

View File

@@ -0,0 +1,390 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [],
"source": [
"import cutlass\n",
"import cutlass.cute as cute"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Tensor\n",
"\n",
"A tensor in CuTe is created through the composition of two key components:\n",
"\n",
"1. An **Engine** (E) - A random-access, pointer-like object that supports:\n",
" - Offset operation: `e + d → e` (offset engine by elements of a layout's codomain)\n",
" - Dereference operation: `*e → v` (dereference engine to produce value)\n",
"\n",
"2. A **Layout** (L) - Defines the mapping from coordinates to offsets\n",
"\n",
"A tensor is formally defined as the composition of an engine E with a layout L, expressed as `T = E ∘ L`. When evaluating a tensor at coordinate c, it:\n",
"\n",
"1. Maps the coordinate c to the codomain using the layout\n",
"2. Offsets the engine accordingly\n",
"3. Dereferences the result to obtain the tensor's value\n",
"\n",
"This can be expressed mathematically as:\n",
"\n",
"```\n",
"T(c) = (E ∘ L)(c) = *(E + L(c))\n",
"```\n",
"\n",
"## Example Usage\n",
"\n",
"Here's a simple example of creating a tensor using pointer and layout `(8,5):(5,1)` and fill with ones:"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [],
"source": [
"@cute.jit\n",
"def create_tensor_from_ptr(ptr: cute.Pointer):\n",
" layout = cute.make_layout((8, 5), stride=(5, 1))\n",
" tensor = cute.make_tensor(ptr, layout)\n",
" tensor.fill(1)\n",
" cute.print_tensor(tensor)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"This creates a tensor where:\n",
"- The engine is a pointer\n",
"- The layout with shape `(8, 5)` and stride `(5, 1)`\n",
"- The resulting tensor can be evaluated using coordinates defined by the layout\n",
"\n",
"We can test this by allocating buffer with torch and run test with pointer to torch tensor"
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor(raw_ptr(0x000000000736b0c0: f32, generic, align<4>) o (8,5):(5,1), data=\n",
" [[ 1.000000, 1.000000, 1.000000, 1.000000, 1.000000, ],\n",
" [ 1.000000, 1.000000, 1.000000, 1.000000, 1.000000, ],\n",
" [ 1.000000, 1.000000, 1.000000, 1.000000, 1.000000, ],\n",
" ...\n",
" [ 1.000000, 1.000000, 1.000000, 1.000000, 1.000000, ],\n",
" [ 1.000000, 1.000000, 1.000000, 1.000000, 1.000000, ],\n",
" [ 1.000000, 1.000000, 1.000000, 1.000000, 1.000000, ]])\n"
]
}
],
"source": [
"import torch\n",
"\n",
"from cutlass.torch import dtype as torch_dtype\n",
"import cutlass.cute.runtime as cute_rt\n",
"\n",
"a = torch.randn(8, 5, dtype=torch_dtype(cutlass.Float32))\n",
"ptr_a = cute_rt.make_ptr(cutlass.Float32, a.data_ptr())\n",
"\n",
"create_tensor_from_ptr(ptr_a)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## DLPACK support \n",
"\n",
"CuTe DSL is designed to support dlpack protocol natively. This offers easy integration with frameworks \n",
"supporting DLPack, e.g. torch, numpy, jax, tensorflow, etc.\n",
"\n",
"For more information, please refer to DLPACK project: https://github.com/dmlc/dlpack\n",
"\n",
"Calling `from_dlpack` can convert any tensor or ndarray object supporting `__dlpack__` and `__dlpack_device__`.\n"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [],
"source": [
"from cutlass.cute.runtime import from_dlpack\n",
"\n",
"@cute.jit\n",
"def print_tensor_dlpack(src: cute.Tensor):\n",
" print(src)\n",
" cute.print_tensor(src)"
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor<ptr<f32, generic> o (8,5):(5,1)>\n",
"tensor(raw_ptr(0x0000000007559340: f32, generic, align<4>) o (8,5):(5,1), data=\n",
" [[-1.151769, 1.019397, -0.371175, -0.717776, 0.502176, ],\n",
" [ 0.114282, 0.900084, 0.320770, 1.564574, -0.632329, ],\n",
" [-0.570140, 0.178112, -0.423079, 1.936198, 0.003355, ],\n",
" ...\n",
" [-2.425393, -0.275528, 1.267157, -0.811101, -0.985456, ],\n",
" [ 0.777889, -2.114074, 0.357184, -0.321312, -0.938138, ],\n",
" [ 1.959564, 1.797602, 0.116901, 0.306198, -1.837295, ]])\n"
]
}
],
"source": [
"a = torch.randn(8, 5, dtype=torch_dtype(cutlass.Float32))\n",
"\n",
"print_tensor_dlpack(from_dlpack(a))"
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor<ptr<f32, generic> o (8,8):(8,1)>\n",
"tensor(raw_ptr(0x0000000007979da0: f32, generic, align<4>) o (8,8):(8,1), data=\n",
" [[ 0.122739, -0.605744, -1.442022, ..., -0.356501, -0.993329, -0.091110, ],\n",
" [ 0.278448, 0.318482, -0.276867, ..., 1.542181, -1.701539, -0.309454, ],\n",
" [ 0.563565, -0.753936, 0.131214, ..., 0.437912, -0.482277, -0.051540, ],\n",
" ...\n",
" [-1.974096, -0.177881, 0.426807, ..., -1.579115, -0.304974, 0.451164, ],\n",
" [ 0.149851, -0.704689, -0.295063, ..., -0.653001, 0.008871, 0.903916, ],\n",
" [ 1.188619, 1.519662, 1.270734, ..., 0.404082, 0.173200, 0.093476, ]])\n"
]
}
],
"source": [
"import numpy as np\n",
"\n",
"a = np.random.randn(8, 8).astype(np.float32)\n",
"\n",
"print_tensor_dlpack(from_dlpack(a))"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Tensor Evaluation Methods\n",
"\n",
"Tensors support two primary methods of evaluation:\n",
"\n",
"### 1. Full Evaluation\n",
"When applying the tensor evaluation with a complete coordinate c, it computes the offset, applies it to the engine, \n",
"and dereferences it to return the stored value. This is the straightforward case where you want to access \n",
"a specific element of the tensor.\n",
"\n",
"### 2. Partial Evaluation (Slicing)\n",
"When evaluating with an incomplete coordinate c = c' ⊕ c* (where c* represents the unspecified portion), \n",
"the result is a new tensor which is a slice of the original tensor with its engine offset to account for \n",
"the coordinates that were provided. This operation can be expressed as:\n",
"\n",
"```\n",
"T(c) = (E ∘ L)(c) = (E + L(c')) ∘ L(c*) = T'(c*)\n",
"```\n",
"\n",
"Slicing effectively reduces the dimensionality of the tensor, creating a sub-tensor that can be \n",
"further evaluated or manipulated."
]
},
{
"cell_type": "code",
"execution_count": 7,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"a[2] = 10.000000 (equivalent to a[(2,0)])\n",
"a[9] = 6.000000 (equivalent to a[(1,1)])\n",
"a[2,0] = 10.000000\n",
"a[2,4] = 14.000000\n",
"a[(2,4)] = 14.000000\n",
"a[2,3] = 100.000000\n",
"a[(2,4)] = 101.000000\n",
"tensor([[ 0., 1., 2., 3., 4.],\n",
" [ 5., 6., 7., 8., 9.],\n",
" [ 10., 11., 12., 100., 101.],\n",
" [ 15., 16., 17., 18., 19.],\n",
" [ 20., 21., 22., 23., 24.],\n",
" [ 25., 26., 27., 28., 29.],\n",
" [ 30., 31., 32., 33., 34.],\n",
" [ 35., 36., 37., 38., 39.]])\n"
]
}
],
"source": [
"@cute.jit\n",
"def tensor_access_item(a: cute.Tensor):\n",
" # access data using linear index\n",
" cute.printf(\"a[2] = {} (equivalent to a[{}])\", a[2],\n",
" cute.make_identity_tensor(a.layout.shape)[2])\n",
" cute.printf(\"a[9] = {} (equivalent to a[{}])\", a[9],\n",
" cute.make_identity_tensor(a.layout.shape)[9])\n",
"\n",
" # access data using n-d coordinates, following two are equivalent\n",
" cute.printf(\"a[2,0] = {}\", a[2, 0])\n",
" cute.printf(\"a[2,4] = {}\", a[2, 4])\n",
" cute.printf(\"a[(2,4)] = {}\", a[2, 4])\n",
"\n",
" # assign value to tensor@(2,4)\n",
" a[2,3] = 100.0\n",
" a[2,4] = 101.0\n",
" cute.printf(\"a[2,3] = {}\", a[2,3])\n",
" cute.printf(\"a[(2,4)] = {}\", a[(2,4)])\n",
"\n",
"@cute.kernel\n",
"def print_tensor_gpu(ptr: cute.Pointer):\n",
" layout = cute.make_layout((8, 5), stride=(5, 1))\n",
" tensor = cute.make_tensor(ptr, layout)\n",
"\n",
" tidx, _, _ = cute.arch.thread_idx()\n",
"\n",
" if tidx == 0:\n",
" cute.print_tensor(tensor)\n",
"\n",
"\n",
"# Create a tensor with sequential data using torch\n",
"data = torch.arange(0, 8*5, dtype=torch.float32).reshape(8, 5)\n",
"tensor_access_item(from_dlpack(data))\n",
"\n",
"print(data)"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Tensor as memory view\n",
"\n",
"In CUDA programming, different memory spaces have different characteristics in terms of access speed, scope, and lifetime:\n",
"\n",
"- **generic**: Default memory space that can refer to any other memory space.\n",
"- **global memory (gmem)**: Accessible by all threads across all blocks, but has higher latency.\n",
"- **shared memory (smem)**: Accessible by all threads within a block, with much lower latency than global memory.\n",
"- **register memory (rmem)**: Thread-private memory with the lowest latency, but limited capacity.\n",
"- **tensor memory (tmem)**: Specialized memory introduced in NVIDIA Blackwell architecture for tensor operations.\n",
"\n",
"When creating tensors in CuTe, you can specify the memory space to optimize performance based on your access patterns.\n",
"\n",
"For more information on CUDA memory spaces, see the [CUDA Programming Guide](https://docs.nvidia.com/cuda/cuda-c-programming-guide/index.html#memory-hierarchy).\n"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Coordinate Tensor\n",
"\n",
"A coordinate tensor is a special type of tensor that maps coordinates to coordinates rather than to values. \n",
"The key distinction is that while regular tensors map coordinates to some value type (like numbers), \n",
"coordinate tensors map coordinates to other coordinates.\n",
"\n",
"For example, given a shape (4,4), a coordinate tensor using row-major layout would appear as:\n",
"\n",
"\\begin{bmatrix} \n",
"(0,0) & (0,1) & (0,2) & (0,3) \\\\\n",
"(1,0) & (1,1) & (1,2) & (1,3) \\\\\n",
"(2,0) & (2,1) & (2,2) & (2,3) \\\\\n",
"(3,0) & (3,1) & (3,2) & (3,3)\n",
"\\end{bmatrix}\n",
"\n",
"The same shape with a column-major layout would appear as:\n",
"\n",
"\\begin{bmatrix}\n",
"(0,0) & (1,0) & (2,0) & (3,0) \\\\\n",
"(0,1) & (1,1) & (2,1) & (3,1) \\\\\n",
"(0,2) & (1,2) & (2,2) & (3,2) \\\\\n",
"(0,3) & (1,3) & (2,3) & (3,3)\n",
"\\end{bmatrix}\n",
"\n",
"The key points about coordinate tensors are:\n",
"- Each element in the tensor is itself a coordinate tuple (i,j) rather than a scalar value\n",
"- The coordinates map to themselves - so position (1,2) contains the coordinate (1,2)\n",
"- The layout (row-major vs column-major) determines how these coordinate tuples are arranged in memory\n",
"\n",
"For example, coordinate tensors can be created using the `make_identity_tensor` utility:\n",
"\n",
"```python\n",
"coord_tensor = make_identity_tensor(layout.shape())\n",
"```\n",
"\n",
"This creates a tensor that maps each coordinate to itself, providing a reference point for understanding how other layouts transform these coordinates."
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor<(0,0) o (8,4):(1@0,1@1)>\n"
]
}
],
"source": [
"@cute.jit\n",
"def print_tensor_coord(a: cute.Tensor):\n",
" coord_tensor = cute.make_identity_tensor(a.layout.shape)\n",
" print(coord_tensor)\n",
"\n",
"a = torch.randn(8,4, dtype=torch_dtype(cutlass.Float32))\n",
"print_tensor_coord(from_dlpack(a))"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.5"
},
"widgets": {
"application/vnd.jupyter.widget-state+json": {
"state": {},
"version_major": 2,
"version_minor": 0
}
}
},
"nbformat": 4,
"nbformat_minor": 4
}

View File

@@ -0,0 +1,558 @@
{
"cells": [
{
"cell_type": "code",
"execution_count": 1,
"metadata": {},
"outputs": [],
"source": [
"import cutlass\n",
"import cutlass.cute as cute\n",
"from cutlass.cute.runtime import from_dlpack\n",
"\n",
"import numpy as np\n",
"import torch"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"# Introduction to the TensorSSA in CuTe DSL\n",
"\n",
"This tutorial introduces what is the `TensorSSA` and why we need it. We also give some examples to show how to use `TensorSSA`.\n",
"\n",
"## What is TensorSSA\n",
"\n",
"`TensorSSA` is a Python class that represents a tensor value in Static Single Assignment (SSA) form within the CuTe DSL. You can think of it as a tensor residing in a (simulated) register.\n",
"\n",
"## Why TensorSSA\n",
"\n",
"`TensorSSA` encapsulates the underlying MLIR tensor value into an object that's easier to manipulate in Python. By overloading numerous Python operators (like `+`, `-`, `*`, `/`, `[]`, etc.), it allows users to express tensor computations (primarily element-wise operations and reductions) in a more Pythonic way. These element-wise operations are then translated into optimized vectorization instructions.\n",
"\n",
"It's part of the CuTe DSL, serving as a bridge between the user-described computational logic and the lower-level MLIR IR, particularly for representing and manipulating register-level data.\n",
"\n",
"## When to use TensorSSA\n",
"\n",
"`TensorSSA` is primarily used in the following scenarios:\n",
"\n",
"### Load from memory and store to memory"
]
},
{
"cell_type": "code",
"execution_count": 2,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"a_vec: tensor_value<vector<12xf32> o (3, 4)>\n",
"b_vec: tensor_value<vector<12xf32> o (3, 4)>\n",
"tensor(raw_ptr(0x0000000006cff170: f32, generic, align<4>) o (3,4):(4,1), data=\n",
" [[ 2.000000, 2.000000, 2.000000, 2.000000, ],\n",
" [ 2.000000, 2.000000, 2.000000, 2.000000, ],\n",
" [ 2.000000, 2.000000, 2.000000, 2.000000, ]])\n"
]
}
],
"source": [
"@cute.jit\n",
"def load_and_store(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):\n",
" \"\"\"\n",
" Load data from memory and store the result to memory.\n",
"\n",
" :param res: The destination tensor to store the result.\n",
" :param a: The source tensor to be loaded.\n",
" :param b: The source tensor to be loaded.\n",
" \"\"\"\n",
" a_vec = a.load()\n",
" print(f\"a_vec: {a_vec}\") # prints `a_vec: vector<12xf32> o (3, 4)`\n",
" b_vec = b.load()\n",
" print(f\"b_vec: {b_vec}\") # prints `b_vec: vector<12xf32> o (3, 4)`\n",
" res.store(a_vec + b_vec)\n",
" cute.print_tensor(res)\n",
"\n",
"a = np.ones(12).reshape((3, 4)).astype(np.float32)\n",
"b = np.ones(12).reshape((3, 4)).astype(np.float32)\n",
"c = np.zeros(12).reshape((3, 4)).astype(np.float32)\n",
"load_and_store(from_dlpack(c), from_dlpack(a), from_dlpack(b))"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"### Register-Level Tensor Operations\n",
"\n",
"When writing kernel logic, various computations, transformations, slicing, etc., are performed on data loaded into registers."
]
},
{
"cell_type": "code",
"execution_count": 3,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor_value<vector<24xf32> o (4, 2, 3)> -> tensor_value<vector<12xf32> o (4, 3)>\n",
"tensor(raw_ptr(0x00000000071acaf0: f32, generic, align<4>) o (4,3):(3,1), data=\n",
" [[ 3.000000, 4.000000, 5.000000, ],\n",
" [ 9.000000, 10.000000, 11.000000, ],\n",
" [ 15.000000, 16.000000, 17.000000, ],\n",
" [ 21.000000, 22.000000, 23.000000, ]])\n"
]
}
],
"source": [
"@cute.jit\n",
"def apply_slice(src: cute.Tensor, dst: cute.Tensor, indices: cutlass.Constexpr):\n",
" \"\"\"\n",
" Apply slice operation on the src tensor and store the result to the dst tensor.\n",
"\n",
" :param src: The source tensor to be sliced.\n",
" :param dst: The destination tensor to store the result.\n",
" :param indices: The indices to slice the source tensor.\n",
" \"\"\"\n",
" src_vec = src.load()\n",
" dst_vec = src_vec[indices]\n",
" print(f\"{src_vec} -> {dst_vec}\")\n",
" if isinstance(dst_vec, cute.TensorSSA):\n",
" dst.store(dst_vec)\n",
" cute.print_tensor(dst)\n",
" else:\n",
" dst[0] = dst_vec\n",
" cute.print_tensor(dst)\n",
"\n",
"def slice_1():\n",
" src_shape = (4, 2, 3)\n",
" dst_shape = (4, 3)\n",
" indices = (None, 1, None)\n",
"\n",
" \"\"\"\n",
" a:\n",
" [[[ 0. 1. 2.]\n",
" [ 3. 4. 5.]]\n",
"\n",
" [[ 6. 7. 8.]\n",
" [ 9. 10. 11.]]\n",
"\n",
" [[12. 13. 14.]\n",
" [15. 16. 17.]]\n",
"\n",
" [[18. 19. 20.]\n",
" [21. 22. 23.]]]\n",
" \"\"\"\n",
" a = np.arange(np.prod(src_shape)).reshape(*src_shape).astype(np.float32)\n",
" dst = np.random.randn(*dst_shape).astype(np.float32)\n",
" apply_slice(from_dlpack(a), from_dlpack(dst), indices)\n",
"\n",
"slice_1()"
]
},
{
"cell_type": "code",
"execution_count": 4,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor_value<vector<24xf32> o (4, 2, 3)> -> ?\n",
"tensor(raw_ptr(0x00000000013cbbe0: f32, generic, align<4>) o (1):(1), data=\n",
" [ 10.000000, ])\n"
]
}
],
"source": [
"def slice_2():\n",
" src_shape = (4, 2, 3)\n",
" dst_shape = (1,)\n",
" indices = 10\n",
" a = np.arange(np.prod(src_shape)).reshape(*src_shape).astype(np.float32)\n",
" dst = np.random.randn(*dst_shape).astype(np.float32)\n",
" apply_slice(from_dlpack(a), from_dlpack(dst), indices)\n",
"\n",
"slice_2()"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"## Arithmetic Operations\n",
"\n",
"As we mentioned earlier, there're many tensor operations whose operands are `TensorSSA`. And they are all element-wise operations. We give some examples below.\n",
"\n",
"### Binary Operations\n",
"\n",
"For binary operations, the LHS operand is `TensorSSA` and the RHS operand can be either `TensorSSA` or `Numeric`. When the RHS is `Numeric`, it will be broadcast to a `TensorSSA`."
]
},
{
"cell_type": "code",
"execution_count": 5,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=\n",
" [ 3.000000, ],\n",
" [ 3.000000, ],\n",
" [ 3.000000, ])\n",
"tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=\n",
" [-1.000000, ],\n",
" [-1.000000, ],\n",
" [-1.000000, ])\n",
"tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=\n",
" [ 2.000000, ],\n",
" [ 2.000000, ],\n",
" [ 2.000000, ])\n",
"tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=\n",
" [ 0.500000, ],\n",
" [ 0.500000, ],\n",
" [ 0.500000, ])\n",
"tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=\n",
" [ 0.000000, ],\n",
" [ 0.000000, ],\n",
" [ 0.000000, ])\n",
"tensor(raw_ptr(0x00000000074f0e70: f32, generic, align<4>) o (3):(1), data=\n",
" [ 1.000000, ],\n",
" [ 1.000000, ],\n",
" [ 1.000000, ])\n"
]
}
],
"source": [
"@cute.jit\n",
"def binary_op_1(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):\n",
" a_vec = a.load()\n",
" b_vec = b.load()\n",
"\n",
" add_res = a_vec + b_vec\n",
" res.store(add_res)\n",
" cute.print_tensor(res) # prints [3.000000, 3.000000, 3.000000]\n",
"\n",
" sub_res = a_vec - b_vec\n",
" res.store(sub_res)\n",
" cute.print_tensor(res) # prints [-1.000000, -1.000000, -1.000000]\n",
"\n",
" mul_res = a_vec * b_vec\n",
" res.store(mul_res)\n",
" cute.print_tensor(res) # prints [2.000000, 2.000000, 2.000000]\n",
"\n",
" div_res = a_vec / b_vec\n",
" res.store(div_res)\n",
" cute.print_tensor(res) # prints [0.500000, 0.500000, 0.500000]\n",
"\n",
" floor_div_res = a_vec // b_vec\n",
" res.store(floor_div_res)\n",
" cute.print_tensor(res) # prints [0.000000, 0.000000, 0.000000]\n",
"\n",
" mod_res = a_vec % b_vec\n",
" res.store(mod_res)\n",
" cute.print_tensor(res) # prints [1.000000, 1.000000, 1.000000]\n",
"\n",
"\n",
"a = np.empty((3,), dtype=np.float32)\n",
"a.fill(1.0)\n",
"b = np.empty((3,), dtype=np.float32)\n",
"b.fill(2.0)\n",
"res = np.empty((3,), dtype=np.float32)\n",
"binary_op_1(from_dlpack(res), from_dlpack(a), from_dlpack(b))"
]
},
{
"cell_type": "code",
"execution_count": 6,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=\n",
" [ 3.000000, ],\n",
" [ 3.000000, ],\n",
" [ 3.000000, ])\n",
"tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=\n",
" [-1.000000, ],\n",
" [-1.000000, ],\n",
" [-1.000000, ])\n",
"tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=\n",
" [ 2.000000, ],\n",
" [ 2.000000, ],\n",
" [ 2.000000, ])\n",
"tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=\n",
" [ 0.500000, ],\n",
" [ 0.500000, ],\n",
" [ 0.500000, ])\n",
"tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=\n",
" [ 0.000000, ],\n",
" [ 0.000000, ],\n",
" [ 0.000000, ])\n",
"tensor(raw_ptr(0x0000000007828ed0: f32, generic, align<4>) o (3):(1), data=\n",
" [ 1.000000, ],\n",
" [ 1.000000, ],\n",
" [ 1.000000, ])\n"
]
}
],
"source": [
"@cute.jit\n",
"def binary_op_2(res: cute.Tensor, a: cute.Tensor, c: cutlass.Constexpr):\n",
" a_vec = a.load()\n",
"\n",
" add_res = a_vec + c\n",
" res.store(add_res)\n",
" cute.print_tensor(res) # prints [3.000000, 3.000000, 3.000000]\n",
"\n",
" sub_res = a_vec - c\n",
" res.store(sub_res)\n",
" cute.print_tensor(res) # prints [-1.000000, -1.000000, -1.000000]\n",
"\n",
" mul_res = a_vec * c\n",
" res.store(mul_res)\n",
" cute.print_tensor(res) # prints [2.000000, 2.000000, 2.000000]\n",
"\n",
" div_res = a_vec / c\n",
" res.store(div_res)\n",
" cute.print_tensor(res) # prints [0.500000, 0.500000, 0.500000]\n",
"\n",
" floor_div_res = a_vec // c\n",
" res.store(floor_div_res)\n",
" cute.print_tensor(res) # prints [0.000000, 0.000000, 0.000000]\n",
"\n",
" mod_res = a_vec % c\n",
" res.store(mod_res)\n",
" cute.print_tensor(res) # prints [1.000000, 1.000000, 1.000000]\n",
"\n",
"a = np.empty((3,), dtype=np.float32)\n",
"a.fill(1.0)\n",
"c = 2.0\n",
"res = np.empty((3,), dtype=np.float32)\n",
"binary_op_2(from_dlpack(res), from_dlpack(a), c)"
]
},
{
"cell_type": "code",
"execution_count": 7,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"[False True False]\n"
]
}
],
"source": [
"@cute.jit\n",
"def binary_op_3(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):\n",
" a_vec = a.load()\n",
" b_vec = b.load()\n",
"\n",
" gt_res = a_vec > b_vec\n",
" res.store(gt_res)\n",
"\n",
" \"\"\"\n",
" ge_res = a_ >= b_ # [False, True, False]\n",
" lt_res = a_ < b_ # [True, False, True]\n",
" le_res = a_ <= b_ # [True, False, True]\n",
" eq_res = a_ == b_ # [False, False, False]\n",
" \"\"\"\n",
"\n",
"a = np.array([1, 2, 3], dtype=np.float32)\n",
"b = np.array([2, 1, 4], dtype=np.float32)\n",
"res = np.empty((3,), dtype=np.bool_)\n",
"binary_op_3(from_dlpack(res), from_dlpack(a), from_dlpack(b))\n",
"print(res) # prints [False, True, False]\n"
]
},
{
"cell_type": "code",
"execution_count": 8,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"[3 0 7]\n"
]
}
],
"source": [
"@cute.jit\n",
"def binary_op_4(res: cute.Tensor, a: cute.Tensor, b: cute.Tensor):\n",
" a_vec = a.load()\n",
" b_vec = b.load()\n",
"\n",
" xor_res = a_vec ^ b_vec\n",
" res.store(xor_res)\n",
"\n",
" # or_res = a_vec | b_vec\n",
" # res.store(or_res) # prints [3, 2, 7]\n",
"\n",
" # and_res = a_vec & b_vec\n",
" # res.store(and_res) # prints [0, 2, 0]\n",
"\n",
"a = np.array([1, 2, 3], dtype=np.int32)\n",
"b = np.array([2, 2, 4], dtype=np.int32)\n",
"res = np.empty((3,), dtype=np.int32)\n",
"binary_op_4(from_dlpack(res), from_dlpack(a), from_dlpack(b))\n",
"print(res) # prints [3, 0, 7]"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"#### Unary Operations"
]
},
{
"cell_type": "code",
"execution_count": 9,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"tensor(raw_ptr(0x0000000007fbd180: f32, generic, align<4>) o (3):(1), data=\n",
" [ 2.000000, ],\n",
" [ 2.000000, ],\n",
" [ 2.000000, ])\n",
"tensor(raw_ptr(0x0000000007fbd180: f32, generic, align<4>) o (3):(1), data=\n",
" [-0.756802, ],\n",
" [-0.756802, ],\n",
" [-0.756802, ])\n",
"tensor(raw_ptr(0x0000000007fbd180: f32, generic, align<4>) o (3):(1), data=\n",
" [ 16.000000, ],\n",
" [ 16.000000, ],\n",
" [ 16.000000, ])\n"
]
}
],
"source": [
"@cute.jit\n",
"def unary_op_1(res: cute.Tensor, a: cute.Tensor):\n",
" a_vec = a.load()\n",
"\n",
" sqrt_res = cute.math.sqrt(a_vec)\n",
" res.store(sqrt_res)\n",
" cute.print_tensor(res) # prints [2.000000, 2.000000, 2.000000]\n",
"\n",
" sin_res = cute.math.sin(a_vec)\n",
" res.store(sin_res)\n",
" cute.print_tensor(res) # prints [-0.756802, -0.756802, -0.756802]\n",
"\n",
" exp2_res = cute.math.exp2(a_vec)\n",
" res.store(exp2_res)\n",
" cute.print_tensor(res) # prints [16.000000, 16.000000, 16.000000]\n",
"\n",
"a = np.array([4.0, 4.0, 4.0], dtype=np.float32)\n",
"res = np.empty((3,), dtype=np.float32)\n",
"unary_op_1(from_dlpack(res), from_dlpack(a))"
]
},
{
"cell_type": "markdown",
"metadata": {},
"source": [
"#### Reduction Operation\n",
"\n",
"The `TensorSSA`'s `reduce` method applies a specified reduction operation (`ReductionOp.ADD`, `ReductionOp.MUL`, `ReductionOp.MAX`, `ReductionOp.MIN`) starting with an initial value, and performs this reduction along the dimensions specified by the `reduction_profile.`. The result is typically a new `TensorSSA` with reduced dimensions or a scalar value if reduces across all axes."
]
},
{
"cell_type": "code",
"execution_count": 10,
"metadata": {},
"outputs": [
{
"name": "stdout",
"output_type": "stream",
"text": [
"21.000000\n",
"tensor(raw_ptr(0x00007ffd1ea2bca0: f32, rmem, align<32>) o (2):(1), data=\n",
" [ 6.000000, ],\n",
" [ 15.000000, ])\n",
"tensor(raw_ptr(0x00007ffd1ea2bcc0: f32, rmem, align<32>) o (3):(1), data=\n",
" [ 6.000000, ],\n",
" [ 8.000000, ],\n",
" [ 10.000000, ])\n"
]
}
],
"source": [
"@cute.jit\n",
"def reduction_op(a: cute.Tensor):\n",
" \"\"\"\n",
" Apply reduction operation on the src tensor.\n",
"\n",
" :param src: The source tensor to be reduced.\n",
" \"\"\"\n",
" a_vec = a.load()\n",
" red_res = a_vec.reduce(\n",
" cute.ReductionOp.ADD,\n",
" 0.0,\n",
" reduction_profile=0\n",
" )\n",
" cute.printf(red_res) # prints 21.000000\n",
"\n",
" red_res = a_vec.reduce(\n",
" cute.ReductionOp.ADD,\n",
" 0.0,\n",
" reduction_profile=(None, 1)\n",
" )\n",
" # We can't print the TensorSSA directly at this point, so we store it to a new Tensor and print it.\n",
" res = cute.make_fragment(red_res.shape, cutlass.Float32)\n",
" res.store(red_res)\n",
" cute.print_tensor(res) # prints [6.000000, 15.000000]\n",
"\n",
" red_res = a_vec.reduce(\n",
" cute.ReductionOp.ADD,\n",
" 1.0,\n",
" reduction_profile=(1, None)\n",
" )\n",
" res = cute.make_fragment(red_res.shape, cutlass.Float32)\n",
" res.store(red_res)\n",
" cute.print_tensor(res) # prints [6.000000, 8.000000, 10.000000]\n",
"\n",
"\n",
"a = np.array([[1, 2, 3], [4, 5, 6]], dtype=np.float32)\n",
"reduction_op(from_dlpack(a))"
]
}
],
"metadata": {
"kernelspec": {
"display_name": "Python 3 (ipykernel)",
"language": "python",
"name": "python3"
},
"language_info": {
"codemirror_mode": {
"name": "ipython",
"version": 3
},
"file_extension": ".py",
"mimetype": "text/x-python",
"name": "python",
"nbconvert_exporter": "python",
"pygments_lexer": "ipython3",
"version": "3.12.5"
}
},
"nbformat": 4,
"nbformat_minor": 4
}