简介
上一章用数学形式描述了 shape 和 stride。把这些整数写成 CuTe 代码时,经常会看到Int<1>{}、_1{}、_128{}。这些到底表示什么呢?这就是本文需要弄懂的知识点。
Int<1>{} 到底是什么?
我们先看一个表达常量的结构体:
template <auto v>
struct C {
using type = C<v>;
static constexpr auto value = v;
using value_type = decltype(v);
CUTE_HOST_DEVICE constexpr operator value_type() const noexcept { return value; }
CUTE_HOST_DEVICE constexpr value_type operator()() const noexcept { return value; }
};
template <int v>
using Int = C<v>;
using _1 = Int<1>;
可以看到,Int只是常量C<v>的别名,也就是说Int<v>表达的是整数常量v。
我们一般可以按照下面的代码使用:
#include <cute/numeric/integral_constant.hpp>
using namespace cute;
int a = 1;
Int<1> b{};
auto c = Int<1>{};
auto d = _1{};
这四个对象的数值都是 1,但 a 的类型是 int,b/c/d 的类型是 Int<1>。
至此,我们可以把 Int<1>{} 拆成三部分:
| 部分 | 含义 |
|---|---|
Int |
CuTe 提供的整数类型模板别名 |
<1> |
模板参数,把数值 1 编码进类型;Int<1> 和 Int<2> 是不同类型 |
{} |
构造一个该类型的对象;其数值由类型中的 1 决定 |
因此,Int<1> 是类型,Int<1>{} 是一个该类型的对象表达式。这里的 {} 是 C++ 初始化语法,不表示空 Tuple,也不是另一个数值参数。
CUTLASS还在库中进一步提供了整数常量类型的别名,例如:
using _1 = Int<1>;
using _2 = Int<2>;
using _128 = Int<128>;
常见的 0、1、2、32、128 等有简写的使用的时候可以直接使用简写;遇到没有定义简写的,直接写 Int<N> 即可。
静态整数与普通整数的区别
在 CuTe 中,静态整数表示“只看类型就知道数值”;普通 int 等类型则承载动态整数。
int n = 128;
auto tile = Int<128>{};
n = 64; // n 仍然是 int,只是保存的值变了
int value = tile; // 可以转换成普通 int,value 的值为 128
auto same_type = tile; // 保留 Int<128> 类型
函数收到一个 int 时,只看类型无法知道传入的是 64 还是 128;收到 Int<128> 时,数值已经包含在类型中。CuTe 因而可以用这个信息推导结果类型、检查约束以及简化计算。
这里说的“动态”是类型表达方式,不代表编译器一定无法优化某次具体计算。
constexpr int 则是另一个容易混淆的地方:
constexpr int k = 128;
auto a = k; // a 的类型是 int
auto b = Int<k>{}; // b 的类型是 Int<128>
k 是 C++ 常量表达式,可以用于数组长度或模板参数;但将它当作普通值传给函数,不会自动变成 CuTe 的静态整数。例如,后面的 make_shape(k) 保存普通整数,make_shape(Int<k>{}) 才把 128 保存在形状的类型中。
Int<N> 的模板参数 N 必须是编译期可用的整数常量。运行时从函数参数或输入中取得的int n 不能直接写成 Int<n>;这种值继续使用普通整数即可。
同样,const int n = ... 的 const 表示不能修改该对象,是否可用于常量表达式还取决于初始化方式。const、constexpr 与“数值编码在类型里”需要分开理解。
模板里写类型,函数参数里传对象
下面两种写法会在后面的代码里反复出现:
using One = Int<1>; // using 后面给出类型
auto one = One{}; // 构造这个类型的对象
看到 Shape<_128, _64> 时,尖括号里传的是类型;看到make_shape(_128{}, _64{}) 时,函数接收的是对象,再从对象推导类型。
单独的 _ 也要和 _1 分开:
_1是整数类型;_1{}表示静态整数 1;_是 CuTe 预定义的占位对象,在后面的 Tensor 切片中用于保留某个维度。
运算之后,静态信息能否保留
以加法为例,相关接口为:
template <auto t, auto u>
CUTE_HOST_DEVICE constexpr C<(t + u)> operator+(C<t>, C<u>);
可以看到,两个静态类型相加的结果依然是一个静态类型。
CuTe 除了为静态类型重载了算数类型操作符,还重载了比较运算符。例如:
int n = 5;
auto a = Int<2>{} + Int<3>{}; // Int<5>
auto b = _2{} * _3{}; // Int<6>
auto c = _2{} + n; // int,值为 7
auto d = _2{} + 3; // int,值为 5
auto e = _0{} * n; // Int<0>,CuTe 对乘以静态零有专门处理
auto f = _2{} < _3{}; // C<true>
两个静态整数做上面的运算,结果继续是静态常量。混入普通整数时,一般会得到普通整数,即便这个普通整数写成字面量 3。某些结果能由静态一侧直接确定,例如 0*n,CuTe 仍然可以保留静态结果。
因此,计算中间值时通常用 auto 保留结果类型:
auto keep = _2{} * _3{}; // 保留 Int<6>
int lose = _2{} * _3{}; // 转成 int,后续类型推导只看到 int
打印并检查类型
本例使用的两个 print 重载分别接收普通整数和静态常量,返回 void:
CUTE_HOST_DEVICE void print(int a);
template <auto Value>
CUTE_HOST_DEVICE void print(C<Value>);
下面是一个完整的入门程序,只在 CPU 上检查常量,不发射 CUDA kernel:
#include <cute/numeric/integral_constant.hpp>
#include <cstdio>
#include <type_traits>
int main() {
using namespace cute;
constexpr int k = 1;
auto ordinary = k;
auto fixed = Int<k>{};
auto sum = _2{} + _3{};
auto mixed = _2{} + 3;
// decltype(x) 取得 x 的类型;is_same_v 比较两个类型是否相同。
// static_assert 在编译期检查条件,不满足就无法编译。
static_assert(std::is_same_v<decltype(ordinary), int>);
static_assert(std::is_same_v<decltype(fixed), _1>);
static_assert(std::is_same_v<decltype(sum), Int<5>>);
static_assert(std::is_same_v<decltype(mixed), int>);
cute::print(ordinary); std::printf("\n");
cute::print(fixed); std::printf("\n");
cute::print(sum); std::printf("\n");
cute::print(mixed); std::printf("\n");
}
按本版本的 print 实现,输出应为:
1
_1
_5
5
输出中的下划线用于区分静态整数和普通整数。_5 与 5 的数值相同,携带的类型信息不同。
Cute中提供了相应的结构体来判断一个数是静态整数还是普通整数:
template <class T>
struct is_static : bool_constant<is_empty<T>::value> {};
template <auto n, class T>
struct is_constant : false_type {};
template <auto n, auto v>
struct is_constant<n, C<v>> : bool_constant<v == n> {};
cute::is_static<T>::value可以检查 CuTe 对静态类型的判定;cute::is_constant<5,T>::value可以进一步检查它是否为值等于 5 的静态整数。