[AMD] Accept ROCm tensors in JIT kernel TensorMatcher + register 4 kernel tests (#29822)

This commit is contained in:
Michael
2026-07-02 19:29:47 -07:00
committed by GitHub
parent a6bc7fef90
commit bee0f34e68
6 changed files with 41 additions and 38 deletions
@@ -62,11 +62,11 @@ void add_constant(tvm::ffi::TensorView dst, tvm::ffi::TensorView src) {
// 1. Validate input tensors
SymbolicSize N = {"num_elements"};
SymbolicDevice device_;
TensorMatcher({N}) // 1D tensor, must be contiguous
.with_dtype<int32_t>() // must be int32
.with_device<kDLCUDA>(device_) // must be on CUDA device
.verify(dst) // check tensor dst
.verify(src); // check tensor src
TensorMatcher({N}) // 1D tensor, must be contiguous
.with_dtype<int32_t>() // must be int32
.with_device<kDLGPU>(device_) // must be on GPU device (CUDA or ROCm)
.verify(dst) // check tensor dst
.verify(src); // check tensor src
// 2. Extract required parameters, prepare for kernel launch
const size_t num_elements = N.unwrap();
@@ -177,19 +177,19 @@ struct CausalConv3dCatPadKernel {
auto out_h = SymbolicSize{"out_h"};
auto out_w = SymbolicSize{"out_w"};
auto device = SymbolicDevice{};
device.set_options<kDLCUDA>();
device.set_options<kDLGPU>();
TensorMatcher({bsz, channels, t_size, h_size, w_size})
.with_dtype<T>()
.template with_device<kDLCUDA>(device)
.template with_device<kDLGPU>(device)
.verify(x);
TensorMatcher({bsz, channels, cache_t, h_size, w_size})
.with_dtype<T>()
.template with_device<kDLCUDA>(device)
.template with_device<kDLGPU>(device)
.verify(cache);
TensorMatcher({bsz, channels, out_t, out_h, out_w})
.with_dtype<T>()
.template with_device<kDLCUDA>(device)
.template with_device<kDLGPU>(device)
.verify(out);
const int64_t depth_left = pad_d_left - cache_t.unwrap();
@@ -209,47 +209,47 @@ struct NgramEmbeddingKernel {
// Verify tensor shapes and types using -1 (kAnySize) for dynamic dimensions
TensorMatcher({-1, -1, -1}) // [ne_n-1, ne_k, ne_n]
.with_dtype<int32_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(ne_weights);
TensorMatcher({-1, -1}) // [ne_n-1, ne_k]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(ne_mods);
TensorMatcher({-1}) // [(ne_n-1)*ne_k + 1]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(exclusive_ne_embeder_size_sums);
TensorMatcher({-1}) // [token_num]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(tokens);
TensorMatcher({-1}) // [batch_size+1]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(exclusive_req_len_sums);
TensorMatcher({-1, -1}) // [max_running_reqs, max_context_len]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(ne_token_table);
TensorMatcher({-1}) // [batch_size]
.with_dtype<int64_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(row_indices);
TensorMatcher({-1}) // [batch_size]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(column_starts);
TensorMatcher({-1, -1}) // [token_num, (ne_n-1)*ne_k]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(n_gram_ids);
const int batch_size = static_cast<int>(exclusive_req_len_sums.size(0) - 1);
@@ -294,37 +294,37 @@ struct NgramEmbeddingKernel {
TensorMatcher({-1, -1, -1}) // [ne_n-1, ne_k, ne_n]
.with_dtype<int32_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(ne_weights);
TensorMatcher({-1, -1}) // [ne_n-1, ne_k]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(ne_mods);
TensorMatcher({-1}) // [(ne_n-1)*ne_k + 1]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(exclusive_ne_embeder_size_sums);
TensorMatcher({-1, -1}) // [max_running_reqs, max_context_len]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(ne_token_table);
TensorMatcher({batch_size}) // [batch_size]
.with_dtype<int64_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(row_indices);
TensorMatcher({batch_size}) // [batch_size]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(column_starts);
TensorMatcher({batch_size, -1}) // [batch_size, (ne_n-1)*ne_k]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(n_gram_ids);
const int bs = static_cast<int>(batch_size.unwrap());
@@ -371,27 +371,27 @@ struct NgramEmbeddingKernel {
// Verify tensor shapes and types using -1 (kAnySize) for dynamic dimensions
TensorMatcher({-1}) // [token_num]
.with_dtype<int32_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(tokens);
TensorMatcher({-1, -1}) // [max_running_reqs, max_context_len]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(ne_token_table);
TensorMatcher({-1}) // [batch_size]
.with_dtype<int64_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(row_indices);
TensorMatcher({-1}) // [batch_size]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(column_starts);
TensorMatcher({-1}) // [batch_size]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(req_lens);
// ignore_tokens can be empty or have values
@@ -400,7 +400,7 @@ struct NgramEmbeddingKernel {
if (has_ignore_tokens) {
TensorMatcher({-1}) // [ignore_token_num]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(ignore_tokens);
}
@@ -447,22 +447,22 @@ struct NgramEmbeddingKernel {
TensorMatcher({batch_size}) // [batch_size]
.with_dtype<int32_t>()
.with_device<kDLCUDA>(device_)
.with_device<kDLGPU>(device_)
.verify(tokens);
TensorMatcher({-1, -1}) // [max_running_reqs, max_context_len]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(ne_token_table);
TensorMatcher({batch_size}) // [batch_size]
.with_dtype<int64_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(row_indices);
TensorMatcher({batch_size}) // [batch_size]
.with_dtype<int32_t>()
.with_device<kDLCUDA>()
.with_device<kDLGPU>()
.verify(column_starts);
const int bs = static_cast<int>(batch_size.unwrap());