Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
23 commits
Select commit Hold shift + click to select a range
c32b424
recognize the patterns used for the RHS matrix
frengels Jan 25, 2022
43931e9
make 1d tile matcher more robust
frengels Jan 25, 2022
f4a153a
put getting rhs tile's index into a separate func
frengels Jan 25, 2022
38a8e96
expand the tests used in correctness check
frengels Jan 25, 2022
cbdc68f
add exclamation mark
frengels Jan 25, 2022
0e3998e
remove unused vars
frengels Jan 26, 2022
464c240
run format and tidy
frengels Jan 26, 2022
c174b98
check for null before using IR in the next step
frengels Jan 28, 2022
f9086fa
check if the broadcast was found
frengels Jan 28, 2022
92426c8
llvm below 13 is no longer supported
frengels Jun 14, 2022
fba5a2e
replace single pattern with commutative permutations
frengels Jun 22, 2022
3764d49
check if the stride is an `IntImm`, otherwise reject pattern
frengels Jun 23, 2022
bd5cd6b
apply clang-format-13
frengels Jun 27, 2022
ddd3a0e
rename wild_i32 -> v2
frengels Jul 15, 2022
db12ffc
check if v1 could be the stride value
frengels Jul 15, 2022
b56aad1
add more detail to a receiving a bad type
frengels Jul 27, 2022
7be1b5b
added short explanation of the right-hand matrix layout
frengels Jul 27, 2022
81bacbb
added explanation for where the 4 comes from
frengels Jul 27, 2022
a80fc76
provide further documentation as to the layout of AMX
frengels Jul 27, 2022
1366699
Merge branch 'main' into pr/6582
steven-johnson Jul 27, 2022
fc89f67
Merge branch 'main' into pr/6582
steven-johnson Jul 27, 2022
314656c
add comments for expected patterns to get_3d_rhs_tile_index
frengels Jul 29, 2022
9b30679
Document the matched pattern
frengels Aug 1, 2022
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
302 changes: 271 additions & 31 deletions src/ExtractTileOperations.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,31 @@
#include "IROperator.h"
#include "Util.h"

/** \file Support extraction of AMX instructions. */

/**
* https://asciiflow.com/#/share/eJyVUkFugzAQ%2FMrKxwoRhdAkza23SmlySHvogQsBp7FkbGSbAoryiz6nr%2BlLugZDk6ghKvJhbXZmd2b3QEScUbIQBece4XFNFVmQQ0SqiCwegtCLSI1RMBtjZGhl8BIRAHh%2BeoFVbBSr4Pq36ZOiSOBpX5cDCEikSGhuipjzun0pmdnD4%2BqtwX9%2Ffg2cLmUcTML76WyO4VAtWJ%2Ff7kIkWMEJ6gbBae2%2F3q53OHBuFBz3TS1HodPqfvUO3%2F4wO7gQag07IXqVkCuZU4VzyApuWI5BAJkdZ0K1B2ZP2%2BwJ%2FEs%2BjhKY0EYViWFSaMAaO6kypBY1hLCtDRIvMTvsekmlsc2kiGgKMw2cxqkGIyEGjn%2FlzonoIMjPUibeQX5Q1bHGisbav%2FBh2kHW2ESzdlaZkqUltaFd9UZ25TnIrIOg%2Bb7vQykLnv661GysRSaSF1k78HkHcaSbntSReLAtTL%2FscOlaI9rxYaRzzgwUOTrZeOCokLzN0TDqRYvUqtFwB6Fvqco9S5r%2BBCiqsWmNLHabzny2Y7E4PyJHcvwBx0t%2BJw%3D%3D)
*
* LHS Matrix RHS Matrix
*
* K conceptually with AMX
* ┌────────┐
* │12345678│ N N*4
*M │ │ ┌──┐ ┌────────┐
* └────────┘ │1 │ K/4│1234 │
* │2 │ │5678 │
* To properly multiply 2 matrices, the │3 │ └────────┘
* AMX instructions perform many 4 byte K│4 │
* dot products, this leads to a lot of │5 │
* striding over 4 byte areas. │6 │
* Normally the row of the LHS matrix, │7 │
* 123... would multiply with the column │8 │
* of the RHS matrix 123..., but with AMX └──┘
* this column is split up into a matrix of columns / 4 byte and rows * 4.
* which then results in K/4 dot products per row.
*
*/

namespace Halide {
namespace Internal {

Expand Down Expand Up @@ -39,12 +64,52 @@ Type amx_op_type_result_type(AMXOpType op_ty) {
}
}

int amx_op_type_size(AMXOpType op_ty) {
switch (op_ty) {
case AMXOpType::Int8:
return 1;
case AMXOpType::Bfloat16:
return 2;
default:
internal_error << "Unexpected";
return -1;
}
}

const auto wild_i32 = Variable::make(Int(32), "*");
const auto wild_i32x = Variable::make(Int(32, 0), "*");

Tile<1> get_1d_tile_index(const Expr &e) {
if (const auto *r1 = e.as<Ramp>()) {
return {true, r1->base, {r1->stride}, {r1->lanes}};

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Do we want some logic akin to HexagonOptimize's apply_commutative_patterns, here and below?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

would that be something like an expr_match that also checks the commutative?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Yes, althought I believe apply_commutative_patterns only checks one level of commutation. I'm just wondering if we should be checking for all of these:

((wild_i32 * wild_i32) + wild_i32) * wild_i32
wild_i32 * ((wild_i32 * wild_i32) + wild_i32)
(wild_i32 + (wild_i32 * wild_i32)) * wild_i32
wild_i32 * (wild_i32 + (wild_i32 * wild_i32))

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

That would make sense, to do that I feel like it requires to implement a new expr_match which applies those rules?

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

For this case, we can just use the four premutations as separate patterns. It might be easier to use the mapped version of expr_match (here).

For the 3d case, we need either explicit patterns or should check if checks on the operands of commutative operations like get_3d_tile_index does here. I'd prefer the former but the latter can be a temporary fast fix for this PR.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

I've played around a bit with extending the expr_match to support commutative operations where relevant, but I think for now I've also settled on the map approach.


const auto stride_var = Variable::make(Int(32), "stride");
const auto v1 = Variable::make(Int(32), "v1");
const auto v2 = Variable::make(Int(32), "v2");
const auto v3 = Variable::make(Int(32), "v3");

Expr patterns[] = {
Comment thread
rootjalex marked this conversation as resolved.
((v1 * stride_var) + v2) * v3,
v3 * ((v1 * stride_var) + v2),
(v2 + (v1 * stride_var)) * v3,
v3 * (v2 + (v1 * stride_var)),
};

std::map<std::string, Expr> matches;
for (const auto &pattern : patterns) {
if (expr_match(pattern, r1->base, matches)) {
auto stride = std::move(matches["stride"]);
// stride must be a constant in order to not be confused with v1
if (stride.as<IntImm>()) {
return {true, r1->base, {std::move(stride)}, {r1->lanes}};
}

// if stride wasn't a constant then v1 could possibly be the stride if constant
auto v1_expr = std::move(matches["v1"]);
if (v1_expr.as<IntImm>()) {
return {true, r1->base, {std::move(v1_expr)}, {r1->lanes}};
}
}
}
}

return {};
Expand Down Expand Up @@ -143,6 +208,169 @@ Tile<3> get_3d_tile_index(const Expr &e) {
return {true, base, {x_stride, 0, r_stride}, {x_tile, y_tile, r_tile}};
}

/**
* \brief Get the 3d rhs tile index configuration
*
* \param e index expression
* \param element_width the width of the elements, 1 for u8/i8, 2 for bf16
* \return Tile<3> the tile configuration found
*
* The pattern which is getting matched looks roughly like
* `broadcast(ramp(0, 1, r), x*y) / broadcast(4, x*y*r) + optional(broadcast(base, x*y*r)) * broadcast(8, x*y*r) +
* broadcast(ramp(0, 1, r), x*y) % broadcast(4, x*y*r) +
* broadcast(ramp(broadcast(_, r), broadcast(4, r), x) , y)`
*/
Tile<3> get_3d_rhs_tile_index(const Expr &e, int element_width) {

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Perhaps not necessary for this PR, but we really should express these as pattern-based rules, they would be far easier to parse and debug

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

Yes, some of it is a bit messy. It would also be nice to have the ability to wildcard on the number of lanes so that arbitrary Ramps and Broadcasts could be matched and have their number of lanes retrieved.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Ah true, that would be quite convenient, perhaps such support could be added to the pattern matcher in a future PR.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

In comment form, could you add a fairly complete list of example Exprs that would match? That will help whoever changes it to pattern matching in the future.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I think this is the last remaining comment request. Once you push a new commit, the buildbots should get themselves unconfused.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

oops, almost didn't see this, should be in 314656c.
I remember the reason why I didn't use patterns for this is because I can't match on a wildcarded number of lanes. And any attempt to use .with_lanes(0) would make asserts in the make functions for the nodes fail since they often require at least 1 lane, but 0 is required to make expr_match ignore the amount of lanes when matching.

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

Sure, but please add example patterns as a comment to help future readers

const auto *sub = e.as<Sub>();
const Add *add_lhs = nullptr;

// there's not always a sub pattern
// This depends on whether we have an ImageParam or a Buffer
if (!sub) {
add_lhs = e.as<Add>();
} else {
add_lhs = sub->a.as<Add>();
}

if (!add_lhs) {
return {};
}

// The right hand side of the add expression is used for retrieving the dimensions of the matrix.
// obtain the x, y, r dimensions
// this expr looks like below, the shape of `add_lhs->a` can be seen further down below
// broadcast(ramp(0, 1, r), x*y) % broadcast(4, x*y*r) + broadcast(ramp(broadcast(base, r), broadcast(4, r), x) , y)
const Add *dim_expr = add_lhs->b.as<Add>();

if (!dim_expr) {
return {};
}

// broadcast(ramp(broadcast(_, r), broadcast(4, r), x), y)
const Broadcast *base_stride_bc = dim_expr->b.as<Broadcast>();

if (!base_stride_bc) {
return {};
}

int tile_y = base_stride_bc->lanes;

// broadcast(ramp(0, 1, r), x*y) % broadcast(4, x*y*r)
const Mod *mod = dim_expr->a.as<Mod>();

if (!mod) {
return {};
}

// broadcast(ramp(0, 1, r), x*y)
const Broadcast *bc_ramp = mod->a.as<Broadcast>();

if (!bc_ramp) {
return {};
}

int tile_xy = bc_ramp->lanes;
int tile_x = tile_xy / tile_y;

// ramp(0, 1, r)
const Ramp *r_ramp = bc_ramp->value.as<Ramp>();

if (!r_ramp) {
return {};
}

int tile_r = r_ramp->lanes;

// get the base and stride
// ramp(broadcast(_, r), broadcast(4, r), x)
const Ramp *base_stride_ramp = base_stride_bc->value.as<Ramp>();

if (!base_stride_ramp) {
return {};
}

// broadcast(_, r)
const Broadcast *base_bc = base_stride_ramp->base.as<Broadcast>();

if (!base_bc) {
return {};
}

Expr base = base_bc->value;
Expr stride;

bool found_stride = false;

// the following pattern will match the following shape
// broadcast(ramp(0, 1, k), x*y) / broadcast(4, x*y*k) * broadcast(_, x*y*k)
// where the stride is marked by _.

// this stride pattern can occur if `tile_r` is the same size as `acc`
auto stride_pattern = Broadcast::make(Ramp::make(0, 1, tile_r), tile_x * tile_y) / Broadcast::make((4 / element_width), tile_x * tile_y * tile_r) * Broadcast::make(wild_i32, tile_x * tile_y * tile_r);

std::vector<Expr> results{};
if (expr_match(stride_pattern, add_lhs->a, results)) {
found_stride = true;
stride = std::move(results[0]);
}

// This pattern is similar to the above except with an additional offset to iterate over the tiles in the k dimension
// (broadcast(ramp(0, 1, k), m * n) / broadcast(4, m*n*k) + _) * broadcast(_, m*n*k)
// here the first _ marks the base and the second _ the stride.
if (!found_stride) {
stride_pattern = (Broadcast::make(Ramp::make(0, 1, tile_r), tile_x * tile_y) / Broadcast::make((4 / element_width), tile_x * tile_y * tile_r) + wild_i32) * Broadcast::make(wild_i32, tile_x * tile_y * tile_r);
if (expr_match(stride_pattern, add_lhs->a, results)) {
found_stride = true;
stride = std::move(results[1]);
base = std::move(results[0]) * stride + base;
}
}

if (!found_stride) {
return {};
}

return {true, base, {stride, 0, 0}, {tile_x, tile_y, tile_r}};
}

struct BaseStride {
bool result{false};
Expr base{};
Expr stride{};
};

BaseStride get_rhs_tile_index(const Expr &index, int element_width, int tile_x, int tile_y, int tile_r) {
const auto rhs_tile2 = get_2d_tile_index(index);

if (!rhs_tile2.result) {
const auto rhs_tile1 = get_1d_tile_index(index);

if (!rhs_tile1.result) {
auto rhs_tile3 = get_3d_rhs_tile_index(index, element_width);
if (rhs_tile3.extent[0] != tile_x || rhs_tile3.extent[1] != tile_y || rhs_tile3.extent[2] != tile_r) {
return {};
}

return {true, rhs_tile3.base, rhs_tile3.stride[0] * element_width};
} else {
if (rhs_tile1.extent[0] != tile_y * tile_r) {
return {};
}

// times 4 because of the rhs layout, each vector used by AMX is 4 bytes in size.
// For the 4 gets divided by the element width which means each vector has 4 elements in u8/i8 and
// 2 elements for bf16.
return {true, rhs_tile1.base, rhs_tile1.stride[0] * (4 / element_width)};
}
} else {
if (tile_y != rhs_tile2.extent[0] || tile_r != rhs_tile2.extent[1]) {
return {};
}

return {true, rhs_tile2.base, rhs_tile2.stride[0]};
}
}

struct Matmul {
bool result = false;
Stmt stmt;
Expand Down Expand Up @@ -197,17 +425,41 @@ Matmul convert_to_matmul(const Store *op, const string &new_name, AMXOpType op_t

const auto *lhs_load = matches[0].as<Load>();
const auto *rhs_broadcast = matches[1].as<Broadcast>();
if (!lhs_load || !rhs_broadcast) {

const Cast *rhs_cast = nullptr;

if (lhs_load && !rhs_broadcast) {
// now working on a larger k dimension

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

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

I don't understand this comment. It seems to describe the PR rather than the code here.

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

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

I've added a slightly larger description here 7be1b5b.
And with further explanation for the layout of the matrices a80fc76.

// with a K dimension of 4 (or 2) with bf16 all the elements in the right-hand matrix are
// layed out in a way that multiplying with a column can be done in a single dot product.
// Therefore the indexing can be reused with a broadcast,
// with higher K dimensions this can no longer be done and the broadcast won't exist.
// ┌──┐
// │1 │
// │2 │
// │3 │ ┌────────┐
// │4 │ │1234 │
// │5 │ │5678 │
// │6 │ └────────┘
// │7 │
// │8 │
// └──┘
rhs_cast = matches[1].as<Cast>();
} else {
rhs_cast = rhs_broadcast->value.as<Cast>();
}

if (!lhs_load || !rhs_cast) {
return {};
}
const auto *rhs_cast = rhs_broadcast->value.as<Cast>();

if (rhs_cast) {
if (op_type == AMXOpType::Int8) {
if (!(rhs_cast->value.type().element_of() == Int(8) || rhs_cast->value.type().element_of() == UInt(8))) {
user_assert(false) << "Expected rhs cast of i8/u8";
}
} else { // AMXOpType::Bfloat16
user_assert(rhs_cast->value.type().element_of() == BFloat(16)) << "Expected rhs cast of bf16";
bool is_i8_u8 = rhs_cast->value.type().element_of() == Int(8) || rhs_cast->value.type().element_of() == UInt(8);
bool is_bf16 = rhs_cast->value.type().element_of() == BFloat(16);

if ((op_type == AMXOpType::Int8 && !is_i8_u8) || (op_type == AMXOpType::Bfloat16 && !is_bf16)) {
user_error << "Expected rhs type of " << (op_type == AMXOpType::Int8 ? "i8/u8" : "bf16")
<< ", got " << rhs_cast->value.type() << " instead.\nIn Expression: " << Expr(rhs_cast);
}
} else {
return {};
Expand All @@ -232,29 +484,15 @@ Matmul convert_to_matmul(const Store *op, const string &new_name, AMXOpType op_t
Expr rhs_base;
Expr rhs_stride;

const auto rhs_tile2 = get_2d_tile_index(rhs_load->index);
if (!rhs_tile2.result) {
const auto rhs_tile1 = get_1d_tile_index(rhs_load->index);

if (!rhs_tile1.result) {
return {};
}

if (rhs_tile1.extent[0] != tile_y * tile_r) {
return {};
}
auto opt_base_stride = get_rhs_tile_index(rhs_load->index, amx_op_type_size(op_type), tile_x, tile_y, tile_r);

rhs_base = rhs_tile1.base;
rhs_stride = rhs_tile1.stride[0];
} else {
if (tile_y != rhs_tile2.extent[0] || tile_r != rhs_tile2.extent[1]) {
return {};
}

rhs_base = rhs_tile2.base;
rhs_stride = rhs_tile2.stride[0];
if (!opt_base_stride.result) {
return {};
}

rhs_base = opt_base_stride.base;
rhs_stride = opt_base_stride.stride;

if (op->index.type().lanes() != tile_x * tile_y ||
factor != tile_r) {
return {};
Expand All @@ -270,7 +508,8 @@ Matmul convert_to_matmul(const Store *op, const string &new_name, AMXOpType op_t
auto rhs_var = Variable::make(Handle(), rhs_load->name);
const auto &rhs_load_type = rhs_load->type;
auto rhs_type = rhs_load_type.with_lanes(1024 / element_width);
auto rhs = Call::make(rhs_type, "tile_load", {1, tile_y * tile_r * element_width, rhs_var, rhs_base * element_width, rhs_stride * tile_y * element_width}, Call::Intrinsic);

auto rhs = Call::make(rhs_type, "tile_load", {tile_r / (4 / element_width), tile_y * 4, rhs_var, rhs_base * element_width, rhs_stride}, Call::Intrinsic);
auto res_type = amx_op_type_result_type(op_type);

// {rows, colbytes, acc, out, lhs, rhs}
Expand Down Expand Up @@ -410,6 +649,7 @@ class ExtractTileOperations : public IRMutator {
found_tile_x = matmul.tile_x;
found_tile_y = matmul.tile_y;
found_tile_r = matmul.tile_r;

return matmul.stmt;
}

Expand All @@ -424,7 +664,7 @@ class ExtractTileOperations : public IRMutator {
}

// Otherwise there is some other operation using the allocation, so we cannot use the AMX instructions
user_assert(false) << "Found non-tile operations for AMX tile allocation";
user_error << "Found non-tile operations for AMX tile allocation";
return op;
}
};
Expand Down
Loading