简介

上一章用数学形式描述了 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 的静态整数。