@@ -1639,22 +1639,16 @@ void CodeGenCUDA::Dispatch_(const AttrStmtNode* op) {
16391639}
16401640
16411641void 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
0 commit comments