From 497742065ea3beb91cc1877ead10e56e3466f8b2 Mon Sep 17 00:00:00 2001 From: Jiacheng Huang Date: Thu, 13 Aug 2026 16:12:00 +0800 Subject: [PATCH] feat(metax): route basic Llama through canonical InfiniOps --- src/infinicore/nn/rope.cc | 2 +- src/infinicore/ops/random_sample/random_sample.cc | 4 +++- .../ops/rotary_embedding/rotary_embedding_infiniops.cc | 5 ++++- xmake.lua | 2 +- 4 files changed, 9 insertions(+), 4 deletions(-) diff --git a/src/infinicore/nn/rope.cc b/src/infinicore/nn/rope.cc index 5e30c7e80..6c7e52ad4 100644 --- a/src/infinicore/nn/rope.cc +++ b/src/infinicore/nn/rope.cc @@ -78,7 +78,7 @@ void RoPE::initialize_cache() { INFINICORE_NN_BUFFER_INIT(cos_cache, ({max_seq_len_, cache_dim}, dtype_, device_)); #ifdef ENABLE_INFINIOPS_API - if (device_.getType() == Device::Type::NVIDIA && !mrope_section_) { + if ((device_.getType() == Device::Type::NVIDIA || device_.getType() == Device::Type::METAX) && !mrope_section_) { INFINICORE_NN_BUFFER_INIT(cos_sin_cache, ({max_seq_len_, rotary_dim_}, dtype_, device_)); } #endif diff --git a/src/infinicore/ops/random_sample/random_sample.cc b/src/infinicore/ops/random_sample/random_sample.cc index 22c9783c0..2ff505c5a 100644 --- a/src/infinicore/ops/random_sample/random_sample.cc +++ b/src/infinicore/ops/random_sample/random_sample.cc @@ -14,7 +14,9 @@ namespace { #ifdef ENABLE_INFINIOPS_API bool tryGreedyWithInfiniOps(Tensor indices, Tensor logits, int topk) { const auto dtype = logits->dtype(); - if (logits->device().getType() != Device::Type::NVIDIA + const auto device_type = logits->device().getType(); + if ((device_type != Device::Type::NVIDIA + && device_type != Device::Type::METAX) || topk != 1 || logits->ndim() != 1 || logits->numel() == 0 diff --git a/src/infinicore/ops/rotary_embedding/rotary_embedding_infiniops.cc b/src/infinicore/ops/rotary_embedding/rotary_embedding_infiniops.cc index 79330f0e0..5c1c5cb19 100644 --- a/src/infinicore/ops/rotary_embedding/rotary_embedding_infiniops.cc +++ b/src/infinicore/ops/rotary_embedding/rotary_embedding_infiniops.cc @@ -31,7 +31,7 @@ void *plan(const Tensor &positions, bool is_neox, int64_t rope_dim_offset, bool inverse) { - INFINICORE_ASSERT(query->device().getType() == Device::Type::NVIDIA); + INFINICORE_ASSERT(query->device().getType() == Device::Type::NVIDIA || query->device().getType() == Device::Type::METAX); return new PlannedMeta{ TensorMeta(positions), TensorMeta(query), @@ -77,6 +77,9 @@ static bool registered = []() { RotaryEmbedding::plan_dispatcher().registerDevice(Device::Type::NVIDIA, &plan); RotaryEmbedding::run_dispatcher().registerDevice(Device::Type::NVIDIA, &run); RotaryEmbedding::cleanup_dispatcher().registerDevice(Device::Type::NVIDIA, &cleanup); + RotaryEmbedding::plan_dispatcher().registerDevice(Device::Type::METAX, &plan); + RotaryEmbedding::run_dispatcher().registerDevice(Device::Type::METAX, &run); + RotaryEmbedding::cleanup_dispatcher().registerDevice(Device::Type::METAX, &cleanup); return true; }(); diff --git a/xmake.lua b/xmake.lua index 7f6c66fb7..b6806dc72 100644 --- a/xmake.lua +++ b/xmake.lua @@ -416,7 +416,7 @@ local function build_infiniops_external(xmake_os) "-DGENERATE_PYTHON_BINDINGS=OFF", "-DCMAKE_BUILD_TYPE=Release" } - if has_config("nv-gpu") then + if has_config("nv-gpu") or has_config("metax-gpu") then table.insert(cmake_config_args, "-DWITH_TORCH=ON") table.insert(cmake_config_args, "-DINFINI_OPS_TORCH_OPS=argmax") end