Skip to content

Commit aa5c0f5

Browse files
committed
Use canonical buffer operation operand names in codegen
1 parent b3c8ce8 commit aa5c0f5

10 files changed

Lines changed: 205 additions & 250 deletions

File tree

‎src/backend/cuda/codegen/codegen_cuda.cc‎

Lines changed: 16 additions & 22 deletions
Original file line numberDiff line numberDiff line change
@@ -1639,22 +1639,16 @@ void CodeGenCUDA::Dispatch_(const AttrStmtNode* op) {
16391639
}
16401640

16411641
void CodeGenCUDA::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) {
1642-
tvm::Tuple allocation_shape = buffer_call->args[0].as_or_throw<tvm::Tuple>();
1643-
auto allocation_extents = allocation_shape->fields.Map(
1642+
auto shape = buffer_call->args[0].as_or_throw<tvm::Tuple>()->fields.Map(
16441643
[](const Expr& extent) { return extent.as_or_throw<PrimExpr>(); });
1645-
DLDataType allocation_dtype_arg = buffer_call->args[1].as_or_throw<DataTypeImm>()->value;
1646-
PrimType allocation_dtype(allocation_dtype_arg);
1647-
ffi::String allocation_scope = buffer_call->args[2].as_or_throw<StringImm>()->value;
1648-
BufferVar allocated_buffer(op->var);
1649-
auto buffer_annotations = buffer_call->attrs.as<DictAttrsNode>()->dict;
1650-
TVM_FFI_ICHECK(allocated_buffer.defined());
1651-
std::string vid = AllocVarID(allocated_buffer.get(), allocated_buffer.name() + "_ptr");
1644+
PrimType dtype(buffer_call->args[1].as_or_throw<DataTypeImm>()->value);
1645+
std::string scope = buffer_call->args[2].as_or_throw<StringImm>()->value;
1646+
BufferVar buffer(op->var);
1647+
auto annotations = buffer_call->attrs.as<DictAttrsNode>()->dict;
1648+
TVM_FFI_ICHECK(buffer.defined());
1649+
std::string vid = AllocVarID(buffer.get(), buffer.name() + "_ptr");
16521650

16531651
this->PrintIndent();
1654-
std::string scope = allocation_scope;
1655-
const VarNode* buffer = allocated_buffer.get();
1656-
PrimType dtype = allocation_dtype;
1657-
16581652
if (scope.find("wmma.") == 0) {
16591653
if (scope == "wmma.matrix_a" || scope == "wmma.matrix_b") {
16601654
bool supported_wmma_input_dtype = dtype == PrimType::Float(16) || dtype == PrimType::Int(8) ||
@@ -1671,12 +1665,12 @@ void CodeGenCUDA::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer
16711665
TVM_FFI_ICHECK(supported_wmma_accumulator_dtype)
16721666
<< "Accumulator only support half, float and int type for now";
16731667
}
1674-
PrintWmmaScope(scope, dtype, buffer, stream);
1668+
PrintWmmaScope(scope, dtype, buffer.get(), stream);
16751669
} else {
16761670
PrintStorageScope(scope, stream);
1677-
int align = allocated_buffer->data_alignment;
1678-
auto it = buffer_annotations.find(tirx::attr::buffer_data_alignment);
1679-
if (it != buffer_annotations.end()) {
1671+
int align = buffer->data_alignment;
1672+
auto it = annotations.find(tirx::attr::buffer_data_alignment);
1673+
if (it != annotations.end()) {
16801674
if (const auto* n = (*it).second.as<IntImmNode>()) {
16811675
align = n->value.as<int>().value();
16821676
}
@@ -1694,15 +1688,15 @@ void CodeGenCUDA::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer
16941688
} else {
16951689
// Compute constant_size from buffer shape
16961690
size_t constant_size = 1;
1697-
for (const auto& dim : allocation_extents) {
1691+
for (const auto& dim : shape) {
16981692
const IntImmNode* dim_imm = dim.as<IntImmNode>();
16991693
TVM_FFI_ICHECK(dim_imm) << "Can only handle constant size stack allocation for now";
17001694
constant_size *= dim_imm->value.as<size_t>().value();
17011695
}
17021696
TVM_FFI_ICHECK_GT(constant_size, 0) << "Can only handle constant size stack allocation for now";
17031697

17041698
if (scope.find("wmma.") == 0) {
1705-
constant_size = GetWmmaFragmentSize(scope, buffer, constant_size);
1699+
constant_size = GetWmmaFragmentSize(scope, buffer.get(), constant_size);
17061700
}
17071701
bool is_packed_integer_dtype =
17081702
dtype == PrimType::Int(4) || dtype == PrimType::UInt(4) || dtype == PrimType::Int(1);
@@ -1712,9 +1706,9 @@ void CodeGenCUDA::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer
17121706
stream << ' ' << vid << '[' << constant_size << "];\n";
17131707
}
17141708

1715-
RegisterHandleType(allocated_buffer.get(), dtype);
1716-
if (buffer_annotations.count(tirx::attr::kVolatile)) {
1717-
MarkVolatile(allocated_buffer.get());
1709+
RegisterHandleType(buffer.get(), dtype);
1710+
if (annotations.count(tirx::attr::kVolatile)) {
1711+
MarkVolatile(buffer.get());
17181712
}
17191713
}
17201714

‎src/backend/cuda/codegen/llvm/codegen_nvptx.cc‎

Lines changed: 13 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -79,29 +79,27 @@ class CodeGenNVPTX : public CodeGenLLVM {
7979
}
8080

8181
void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) final {
82-
tvm::Tuple allocation_shape = buffer_call->args[0].as_or_throw<tvm::Tuple>();
83-
Array<PrimExpr> allocation_extents = allocation_shape->fields.as_or_throw<Array<PrimExpr>>();
84-
DLDataType allocation_dtype_arg = buffer_call->args[1].as_or_throw<DataTypeImm>()->value;
85-
PrimType allocation_dtype(allocation_dtype_arg);
86-
ffi::String allocation_scope = buffer_call->args[2].as_or_throw<StringImm>()->value;
87-
BufferVar allocated_buffer(op->var);
88-
auto buffer_annotations = buffer_call->attrs.as<DictAttrsNode>()->dict;
82+
Array<PrimExpr> shape =
83+
buffer_call->args[0].as_or_throw<tvm::Tuple>()->fields.as_or_throw<Array<PrimExpr>>();
84+
PrimType dtype(buffer_call->args[1].as_or_throw<DataTypeImm>()->value);
85+
ffi::String scope = buffer_call->args[2].as_or_throw<StringImm>()->value;
86+
BufferVar buffer(op->var);
87+
auto annotations = buffer_call->attrs.as<DictAttrsNode>()->dict;
8988
llvm::Value* buf = nullptr;
90-
StorageInfo& info = alloc_storage_info_[allocated_buffer.get()];
89+
StorageInfo& info = alloc_storage_info_[buffer.get()];
9190
// maximum necessary alignment in the NV devices
9291
if (info.alignment > 16) {
9392
info.alignment = 16;
9493
}
9594

96-
auto storage_scope = runtime::StorageScope::Create(allocation_scope);
97-
PrimType dtype = allocation_dtype;
95+
auto storage_scope = runtime::StorageScope::Create(scope);
9896

9997
if (storage_scope.rank == runtime::StorageRank::kShared && storage_scope.tag == ".dyn") {
10098
// Shared memory: address space == 3
10199
buf = AllocateSharedMemory(dtype, 0, 3, info.alignment, llvm::GlobalValue::ExternalLinkage);
102100
} else {
103101
// Compute constant_size from buffer shape
104-
const IntImmNode* dim_imm = allocation_extents[0].as<IntImmNode>();
102+
const IntImmNode* dim_imm = shape[0].as<IntImmNode>();
105103
TVM_FFI_ICHECK(dim_imm) << "Can only handle constant size stack allocation in GPU";
106104
size_t constant_size = dim_imm->value.as<size_t>().value();
107105
TVM_FFI_ICHECK_GT(constant_size, 0)
@@ -129,10 +127,10 @@ class CodeGenNVPTX : public CodeGenLLVM {
129127

130128
buf = builder_->CreatePointerCast(
131129
buf, llvmGetPointerTo(DTypeToLLVMType(dtype), buf->getType()->getPointerAddressSpace()));
132-
TVM_FFI_ICHECK(!var_map_.count(allocated_buffer.get()));
133-
var_map_[allocated_buffer.get()] = buf;
134-
if (buffer_annotations.count(tirx::attr::kVolatile)) {
135-
volatile_buf_.insert(allocated_buffer.get());
130+
TVM_FFI_ICHECK(!var_map_.count(buffer.get()));
131+
var_map_[buffer.get()] = buf;
132+
if (annotations.count(tirx::attr::kVolatile)) {
133+
volatile_buf_.insert(buffer.get());
136134
}
137135
}
138136

‎src/backend/metal/codegen/codegen_metal.cc‎

Lines changed: 13 additions & 17 deletions
Original file line numberDiff line numberDiff line change
@@ -366,21 +366,19 @@ void CodeGenMetal::Dispatch_(const BindNode* op) {
366366
}
367367

368368
void CodeGenMetal::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) {
369-
tvm::Tuple allocation_shape = buffer_call->args[0].as_or_throw<tvm::Tuple>();
370-
auto allocation_extents = allocation_shape->fields.Map(
369+
auto shape = buffer_call->args[0].as_or_throw<tvm::Tuple>()->fields.Map(
371370
[](const Expr& extent) { return extent.as_or_throw<PrimExpr>(); });
372-
DLDataType allocation_dtype_arg = buffer_call->args[1].as_or_throw<DataTypeImm>()->value;
373-
PrimType allocation_dtype(allocation_dtype_arg);
374-
ffi::String allocation_scope = buffer_call->args[2].as_or_throw<StringImm>()->value;
375-
BufferVar allocated_buffer(op->var);
376-
auto buffer_annotations = buffer_call->attrs.as<DictAttrsNode>()->dict;
377-
TVM_FFI_ICHECK(allocated_buffer.defined());
378-
std::string vid = AllocVarID(allocated_buffer.get());
371+
PrimType dtype(buffer_call->args[1].as_or_throw<DataTypeImm>()->value);
372+
std::string scope = buffer_call->args[2].as_or_throw<StringImm>()->value;
373+
BufferVar buffer(op->var);
374+
auto annotations = buffer_call->attrs.as<DictAttrsNode>()->dict;
375+
TVM_FFI_ICHECK(buffer.defined());
376+
std::string vid = AllocVarID(buffer.get());
379377

380378
this->PrintIndent();
381379
// Compute a compile-time upper bound on the number of buffer elements.
382380
size_t constant_size = 1;
383-
for (const auto& dim : allocation_extents) {
381+
for (const auto& dim : shape) {
384382
const auto* dim_imm = dim.as<IntImmNode>();
385383
int64_t dim_size =
386384
dim_imm ? static_cast<int64_t>(dim_imm->value) : analyzer_->const_int_bound(dim)->max_value;
@@ -402,9 +400,7 @@ void CodeGenMetal::DispatchAllocBuffer(const BindNode* op, const CallNode* buffe
402400
constant_size *= static_cast<size_t>(dim_size);
403401
}
404402

405-
auto scope = allocation_scope;
406-
alloc_storage_scope_[allocated_buffer.get()] = scope;
407-
const PrimType& dtype = allocation_dtype;
403+
alloc_storage_scope_[buffer.get()] = scope;
408404
if (scope == "metal.simdgroup") {
409405
bool supported_simdgroup_dtype = dtype == PrimType::Float(16) || dtype == PrimType::Float(32) ||
410406
dtype == PrimType::BFloat(16);
@@ -417,17 +413,17 @@ void CodeGenMetal::DispatchAllocBuffer(const BindNode* op, const CallNode* buffe
417413
std::ostringstream dtype_os;
418414
PrintType(dtype, dtype_os);
419415
std::string dtype_str = dtype_os.str();
420-
simdgroup_dtype_[allocated_buffer.get()] = dtype_str;
416+
simdgroup_dtype_[buffer.get()] = dtype_str;
421417
stream << "simdgroup_" << dtype_str << "8x8 " << vid << '[' << constant_size / 64 << "];\n";
422418
} else {
423419
PrintStorageScope(scope, stream);
424420
PrintType(dtype, stream);
425421
stream << ' ' << vid << '[' << constant_size << "];\n";
426422
}
427423

428-
RegisterHandleType(allocated_buffer.get(), allocation_dtype);
429-
if (buffer_annotations.count(tirx::attr::kVolatile)) {
430-
MarkVolatile(allocated_buffer.get());
424+
RegisterHandleType(buffer.get(), dtype);
425+
if (annotations.count(tirx::attr::kVolatile)) {
426+
MarkVolatile(buffer.get());
431427
}
432428
}
433429

‎src/backend/opencl/codegen/codegen_opencl.cc‎

Lines changed: 5 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -449,20 +449,18 @@ std::string CodeGenOpenCL::CastTo(std::string value, const PrimType& target) {
449449
}
450450

451451
void CodeGenOpenCL::DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) {
452-
tvm::Tuple allocation_shape = buffer_call->args[0].as_or_throw<tvm::Tuple>();
453-
auto allocation_extents = allocation_shape->fields.Map(
452+
auto shape = buffer_call->args[0].as_or_throw<tvm::Tuple>()->fields.Map(
454453
[](const Expr& extent) { return extent.as_or_throw<PrimExpr>(); });
455-
DLDataType allocation_dtype_arg = buffer_call->args[1].as_or_throw<DataTypeImm>()->value;
456-
PrimType allocation_dtype(allocation_dtype_arg);
457-
BufferVar allocated_buffer(op->var);
454+
PrimType dtype(buffer_call->args[1].as_or_throw<DataTypeImm>()->value);
455+
BufferVar buffer(op->var);
458456
// Compute constant_size from buffer shape
459457
size_t constant_size = 1;
460-
for (const auto& dim : allocation_extents) {
458+
for (const auto& dim : shape) {
461459
const IntImmNode* dim_imm = dim.as<IntImmNode>();
462460
TVM_FFI_ICHECK(dim_imm) << "Can only handle constant size stack allocation for now";
463461
constant_size *= dim_imm->value.as<size_t>().value();
464462
}
465-
allocation_size_.insert({allocated_buffer.get(), constant_size * allocation_dtype.lanes()});
463+
allocation_size_.insert({buffer.get(), constant_size * dtype.lanes()});
466464
CodeGenC::DispatchAllocBuffer(op, buffer_call);
467465
}
468466

‎src/backend/rocm/codegen/llvm/codegen_amdgpu.cc‎

Lines changed: 12 additions & 15 deletions
Original file line numberDiff line numberDiff line change
@@ -97,25 +97,22 @@ class CodeGenAMDGPU : public CodeGenLLVM {
9797
}
9898

9999
void DispatchAllocBuffer(const BindNode* op, const CallNode* buffer_call) final {
100-
tvm::Tuple allocation_shape = buffer_call->args[0].as_or_throw<tvm::Tuple>();
101-
auto allocation_extents = allocation_shape->fields.Map(
100+
auto shape = buffer_call->args[0].as_or_throw<tvm::Tuple>()->fields.Map(
102101
[](const Expr& extent) { return extent.as_or_throw<PrimExpr>(); });
103-
DLDataType allocation_dtype_arg = buffer_call->args[1].as_or_throw<DataTypeImm>()->value;
104-
PrimType allocation_dtype(allocation_dtype_arg);
105-
ffi::String allocation_scope = buffer_call->args[2].as_or_throw<StringImm>()->value;
106-
BufferVar allocated_buffer(op->var);
107-
auto buffer_annotations = buffer_call->attrs.as<DictAttrsNode>()->dict;
102+
PrimType dtype(buffer_call->args[1].as_or_throw<DataTypeImm>()->value);
103+
ffi::String scope = buffer_call->args[2].as_or_throw<StringImm>()->value;
104+
BufferVar buffer(op->var);
105+
auto annotations = buffer_call->attrs.as<DictAttrsNode>()->dict;
108106
llvm::Value* buf = nullptr;
109-
StorageInfo& info = alloc_storage_info_[allocated_buffer.get()];
110-
auto storage_scope = runtime::StorageScope::Create(allocation_scope);
111-
PrimType dtype = allocation_dtype;
107+
StorageInfo& info = alloc_storage_info_[buffer.get()];
108+
auto storage_scope = runtime::StorageScope::Create(scope);
112109

113110
if (storage_scope.rank == runtime::StorageRank::kShared && storage_scope.tag == ".dyn") {
114111
LOG(WARNING) << "Dynamic shared memory support for rocm is experimental.";
115112
buf = AllocateSharedMemory(dtype, 0, 3, std::min(info.alignment, 16),
116113
llvm::GlobalValue::ExternalLinkage);
117114
} else {
118-
const IntImmNode* dim_imm = allocation_extents[0].as<IntImmNode>();
115+
const IntImmNode* dim_imm = shape[0].as<IntImmNode>();
119116
TVM_FFI_ICHECK(dim_imm) << "Can only handle constant size stack allocation in GPU";
120117
size_t constant_size = dim_imm->value.as<size_t>().value();
121118
TVM_FFI_ICHECK_GT(constant_size, 0)
@@ -148,10 +145,10 @@ class CodeGenAMDGPU : public CodeGenLLVM {
148145

149146
buf = builder_->CreatePointerCast(
150147
buf, llvmGetPointerTo(DTypeToLLVMType(dtype), buf->getType()->getPointerAddressSpace()));
151-
TVM_FFI_ICHECK(!var_map_.count(allocated_buffer.get()));
152-
var_map_[allocated_buffer.get()] = buf;
153-
if (buffer_annotations.count(tirx::attr::kVolatile)) {
154-
volatile_buf_.insert(allocated_buffer.get());
148+
TVM_FFI_ICHECK(!var_map_.count(buffer.get()));
149+
var_map_[buffer.get()] = buf;
150+
if (annotations.count(tirx::attr::kVolatile)) {
151+
volatile_buf_.insert(buffer.get());
155152
}
156153
}
157154

0 commit comments

Comments
 (0)