Files
Tianqi Chen 1e1920bcbd [REFACTOR][IR] Unify PrimExpr type mechanism to PrimType instead of DataType (#19875)
In the past we have been using `DataType` in PrimExpr.dtype field to
check type information for PrimExpr while still having BaseExpr.ty for
richer type information. DataType is also used both in runtime and
compiler. This PR streamlines the boundary:

- PrimExpr.ty now carries PrimType that replaces original use of
`DataType`
- Runtime use will now favor DLPack DLDataType, removing one layer of
indirection.
- Constants attributes where values are usually runtime values, will use
`DLDataType`
- DataType will be phased out after this PR

We also brings up helper functions in PrimType, but also limits them to
a more concise set so the functions do not grow with the data type codes
in DLPack.

This is a major refactor that changes the IR primitive. It helps to
bring possible future benefits:
- Unified type mechanism through Expr.ty
- Possibility of carry future Type nodes 

Migration Guide:
- Use `PrimType` when code reasons about compiler expression types,
tensor element compiler types, or constructs a `PrimExpr`/compiler type.
- Use existing source types such as `expr.ty()`, `ExprOp.expr_ty()`, or
TE tensor element `dtype` where possible instead of rebuilding a type
from dtype text.
- Use raw `DLDataType` for runtime constants, ABI paths, dtype-valued
attrs, and storage/runtime helper logic.
- Prefer direct `PrimType` equality, `MatchesCode(...)`,
`MatchesElementType(...)`, and `WithCode(...)` over local wrappers or
string dtype checks.

Performance:

Using Object type instead of DLDataType would indeed bring some
performance impact to the IR. We have done the following performance
optimizations:
- Make sure most of the outputs reuse one of the PrimType from inputs
- Cache a thread local PrimType based on input so we don't repeatly
realloc

We did benchmarks show that rewrite simplify operation stays within
+-10% overhead of original one. Which merits the refactor given the
benefit the unfication brings
2026-06-24 21:31:47 -04:00

74 lines
2.5 KiB
C++

/*
* Licensed to the Apache Software Foundation (ASF) under one
* or more contributor license agreements. See the NOTICE file
* distributed with this work for additional information
* regarding copyright ownership. The ASF licenses this file
* to you under the Apache License, Version 2.0 (the
* "License"); you may not use this file except in compliance
* with the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing,
* software distributed under the License is distributed on an
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
* KIND, either express or implied. See the License for the
* specific language governing permissions and limitations
* under the License.
*/
#include <gtest/gtest.h>
#include <tvm/runtime/logging.h>
#include <tvm/runtime/tensor.h>
using namespace tvm;
TEST(TensorTest, IsContiguous_ContiguousStride) {
auto array = runtime::Tensor::Empty({5, 10}, DLDataType{kDLFloat, 32, 1}, {kDLCPU});
DLManagedTensor* managed_tensor = array.ToDLPack();
int64_t strides[] = {10, 1};
managed_tensor->dl_tensor.strides = strides;
TVM_FFI_ICHECK(ffi::IsContiguous(managed_tensor->dl_tensor));
managed_tensor->deleter(managed_tensor);
}
TEST(TensorTest, IsContiguous_NullStride) {
auto array = runtime::Tensor::Empty({5, 10}, DLDataType{kDLFloat, 32, 1}, {kDLCPU});
DLManagedTensor* managed_tensor = array.ToDLPack();
managed_tensor->dl_tensor.strides = nullptr;
TVM_FFI_ICHECK(ffi::IsContiguous(managed_tensor->dl_tensor));
managed_tensor->deleter(managed_tensor);
}
TEST(TensorTest, IsContiguous_AnyStrideForSingular) {
auto array = runtime::Tensor::Empty({5, 1, 10}, DLDataType{kDLFloat, 32, 1}, {kDLCPU});
DLManagedTensor* managed_tensor = array.ToDLPack();
int64_t strides[] = {10, 1, 1}; // strides[1] is normalized to 1 because shape[1] == 1.
managed_tensor->dl_tensor.strides = strides;
TVM_FFI_ICHECK(ffi::IsContiguous(managed_tensor->dl_tensor));
managed_tensor->dl_tensor.strides = nullptr;
managed_tensor->deleter(managed_tensor);
}
TEST(TensorTest, IsContiguous_UncontiguousStride) {
auto array = runtime::Tensor::Empty({5, 1, 10}, DLDataType{kDLFloat, 32, 1}, {kDLCPU});
DLManagedTensor* managed_tensor = array.ToDLPack();
int64_t strides[] = {1, 1, 1};
managed_tensor->dl_tensor.strides = strides;
TVM_FFI_ICHECK(!ffi::IsContiguous(managed_tensor->dl_tensor));
managed_tensor->dl_tensor.strides = nullptr;
managed_tensor->deleter(managed_tensor);
}