From 0c2b716ac5a29a065ffaeeb8c4c68cb0e31690e4 Mon Sep 17 00:00:00 2001 From: cx Date: Thu, 27 Aug 2026 05:35:22 +0000 Subject: [PATCH] fix: address backend-independent correctness issues - honor DDP reduction configuration and validate process group devices - fix elementwise and CUDA linear backward correctness - remove the unsupported distributed optimizer test - clean up related build diagnostics and test documentation --- CMakeLists.txt | 15 +++-- docs/test_infrastructure_design.md | 11 ++-- docs/test_usage_guide.md | 8 ++- example/gpt2/main.cc | 1 + example/llama3/main.cc | 1 + infini_train/include/autograd/elementwise.h | 2 + .../include/nn/parallel/process_group.h | 2 + infini_train/src/autograd/elementwise.cc | 18 +++--- infini_train/src/core/runtime/device_guard.cc | 2 +- infini_train/src/kernels/cuda/linear.cu | 21 +++---- .../parallel/ddp/distributed_data_parallel.cc | 25 ++++++--- infini_train/src/nn/parallel/ddp/reducer.cc | 4 +- infini_train/src/tensor.cc | 2 + .../test_autograd_elementwise_backward.cc | 55 +++++++++++++++++-- .../autograd/test_autograd_linear_backward.cc | 14 ++++- tests/optimizer/CMakeLists.txt | 14 ----- .../test_optimizer_parameter_names.cc | 40 -------------- 17 files changed, 130 insertions(+), 105 deletions(-) diff --git a/CMakeLists.txt b/CMakeLists.txt index 6bd8069d..4ffbc25e 100644 --- a/CMakeLists.txt +++ b/CMakeLists.txt @@ -3,7 +3,7 @@ cmake_minimum_required(VERSION 3.28) option(USE_CUDA "Support NVIDIA CUDA" OFF) option(PROFILE_MODE "ENABLE PROFILE MODE" OFF) option(USE_OMP "Use OpenMP as backend for Eigen" ON) -option(USE_NCCL "Build project for distributed running" ON) +option(USE_NCCL "Build project for distributed running on CUDA using NCCL" ON) option(BUILD_TEST "Build InfiniTrain tests" OFF) project(infini_train VERSION 0.6.0 LANGUAGES CXX) @@ -64,12 +64,15 @@ endif() # Framework core sources (*.cc), excluding cpu kernels (they are built separately) file(GLOB_RECURSE SRC ${PROJECT_SOURCE_DIR}/infini_train/src/*.cc) list(FILTER SRC EXCLUDE REGEX ".*kernels/cpu/.*") + +# Exclude backend-specific runtime/ccl translation units when the corresponding +# backend is disabled. This keeps each build self-contained and avoids pulling +# in headers (e.g. / ) that are not on the +# include path. if(NOT USE_CUDA) - list(FILTER SRC EXCLUDE REGEX ".*runtime/cuda/.*") - list(FILTER SRC EXCLUDE REGEX ".*ccl/cuda/.*") -endif() -if(NOT USE_NCCL) - list(FILTER SRC EXCLUDE REGEX ".*infini_train/src/core/ccl/cuda/.*") + list(FILTER SRC EXCLUDE REGEX ".*/(ccl|runtime)/cuda/.*") +elseif(NOT USE_NCCL) + list(FILTER SRC EXCLUDE REGEX ".*/ccl/cuda/.*") endif() # CPU kernels (*.cc) diff --git a/docs/test_infrastructure_design.md b/docs/test_infrastructure_design.md index 8aa210ce..857824e1 100644 --- a/docs/test_infrastructure_design.md +++ b/docs/test_infrastructure_design.md @@ -8,7 +8,7 @@ tests/ ├── CMakeLists.txt # 顶层:include 宏 + add_subdirectory ├── common/ -│ ├── CMakeLists.txt # header-only interface library +│ ├── CMakeLists.txt # 公共 test_main target │ └── test_utils.h # C++ 基类、skip 宏、填充工具函数 ├── tensor/ # Tensor 创建 / 拷贝 / 销毁 / 算子 ├── optimizer/ # Optimizer 创建 / step @@ -16,7 +16,8 @@ tests/ ├── hook/ # Module hook + precision check ├── lora/ # LoRA 相关 ├── dtype/ # Scalar / dtype dispatch + 编译期负面测试 -└── transformer/ # Transformer 架构测试 +├── transformer/ # Transformer 架构测试 +└── checkpoint/ # Checkpoint 序列化测试 cmake/ └── test_macros.cmake # CMake 宏:infini_train_add_test / infini_train_add_test_suite @@ -88,11 +89,11 @@ ctest -L cpu --output-on-failure ctest -L cuda --output-on-failure # 运行单个测试二进制(看完整 GTest 输出) -./test_tensor_cpu -./test_autograd_cuda +./tests/tensor/test_tensor_cpu +./tests/autograd/test_autograd_cuda # GTest filter 过滤特定用例 -./test_tensor_cpu --gtest_filter="CPU/TensorCreateTest.*" +./tests/tensor/test_tensor_cpu --gtest_filter="CPU/TensorCreateTest.*" ``` 无 GPU 机器上 `cmake -DBUILD_TEST=ON -DUSE_CUDA=OFF ..` 即可,CUDA 测试实例不会注册。 diff --git a/docs/test_usage_guide.md b/docs/test_usage_guide.md index 45113e52..248f1a3f 100644 --- a/docs/test_usage_guide.md +++ b/docs/test_usage_guide.md @@ -36,8 +36,8 @@ ctest -L cuda --output-on-failure ctest -R tensor --output-on-failure # 直接运行测试二进制,使用 GTest 过滤器 -./tests/tensor/test_tensor_create_cpu --gtest_filter="CPU/TensorCreateTest.*" -./tests/tensor/test_tensor_create_cuda --gtest_filter="CUDA/TensorCreateTest.*" +./tests/tensor/test_tensor_cpu --gtest_filter="CPU/TensorCreateTest.*" +./tests/tensor/test_tensor_cuda --gtest_filter="CUDA/TensorCreateTest.*" ``` --- @@ -75,7 +75,9 @@ INFINI_TRAIN_REGISTER_TEST(TensorCopyTest); 在子目录的 `CMakeLists.txt`(例如 `tests/tensor/CMakeLists.txt`)中添加: ```cmake -infini_train_add_test_suite(test_tensor_copy test_tensor_copy.cc) +infini_train_add_test_suite(test_tensor_copy + SOURCES test_tensor_copy.cc +) ``` 这会生成两个 CTest 目标:`test_tensor_copy_cpu`(标签 `cpu`)和 `test_tensor_copy_cuda`(标签 `cuda`)。 diff --git a/example/gpt2/main.cc b/example/gpt2/main.cc index 2551880e..cb8c2e26 100644 --- a/example/gpt2/main.cc +++ b/example/gpt2/main.cc @@ -4,6 +4,7 @@ #include #include #include +#include #include #include diff --git a/example/llama3/main.cc b/example/llama3/main.cc index ccfca86a..c0a03958 100644 --- a/example/llama3/main.cc +++ b/example/llama3/main.cc @@ -3,6 +3,7 @@ #include #include #include +#include #include #include "gflags/gflags.h" diff --git a/infini_train/include/autograd/elementwise.h b/infini_train/include/autograd/elementwise.h index c4333b64..16c33c4f 100644 --- a/infini_train/include/autograd/elementwise.h +++ b/infini_train/include/autograd/elementwise.h @@ -104,6 +104,8 @@ class Exp : public Function { explicit Exp() : Function(kType) {} std::vector> Forward(const std::vector> &input_tensors) override; + void SetupContext(const std::vector> &input_tensors, + const std::vector> &output_tensors) override; std::vector> Backward(const std::vector> &grad_outputs) override; }; diff --git a/infini_train/include/nn/parallel/process_group.h b/infini_train/include/nn/parallel/process_group.h index 2c439f40..ef27e67f 100644 --- a/infini_train/include/nn/parallel/process_group.h +++ b/infini_train/include/nn/parallel/process_group.h @@ -39,6 +39,8 @@ class ProcessGroup { virtual int GetGroupRank(int global_rank) const; + Device::DeviceType backend() const { return backend_; } + // Asynchronous communication APIs (Compute / Communication stream decoupled) virtual std::shared_ptr AllReduce(const std::shared_ptr &tensor, function::ReduceOpType reduce_op = function::ReduceOpType::kSum, diff --git a/infini_train/src/autograd/elementwise.cc b/infini_train/src/autograd/elementwise.cc index 36a2cad7..f9403d8d 100644 --- a/infini_train/src/autograd/elementwise.cc +++ b/infini_train/src/autograd/elementwise.cc @@ -182,12 +182,21 @@ std::vector> Exp::Forward(const std::vector>({device, "ExpForward"}, input)}; } +void Exp::SetupContext(const std::vector> &, + const std::vector> &output_tensors) { + const auto &output = output_tensors[0]; + ctx_.SaveForBackward({output}); +} + std::vector> Exp::Backward(const std::vector> &grad_outputs) { + auto saved_tensors = ctx_.GetSavedTensors(); + CHECK_EQ(saved_tensors.size(), 1); + const auto &output = saved_tensors[0]; CHECK_EQ(grad_outputs.size(), 1); const auto &grad_output = grad_outputs[0]; - auto device = grad_output->GetDevice().type(); - return {Dispatcher::Instance().Call>({device, "ExpBackward"}, grad_output)}; + auto device = output->GetDevice().type(); + return {Dispatcher::Instance().Call>({device, "ExpBackward"}, grad_output, output)}; } std::vector> Log::Forward(const std::vector> &input_tensors) { @@ -397,11 +406,6 @@ std::vector> Add::Backward(const std::vectorGetDevice().type(); auto [grad_a, grad_b] = Dispatcher::Instance().Call, std::shared_ptr>>( {device, "AddBackward"}, grad_output, a_dims_, b_dims_); diff --git a/infini_train/src/core/runtime/device_guard.cc b/infini_train/src/core/runtime/device_guard.cc index fbcb316f..d57b0d7c 100644 --- a/infini_train/src/core/runtime/device_guard.cc +++ b/infini_train/src/core/runtime/device_guard.cc @@ -140,7 +140,7 @@ void DeviceGuardImplRegistry::Register(Device::DeviceType type, std::unique_ptr< } if (impls_.contains(type)) { - LOG(FATAL) << std::format("DeviceGuardImpl for type {} already registrered", static_cast(type)); + LOG(FATAL) << std::format("DeviceGuardImpl for type {} already registered", static_cast(type)); } if (!impls_.empty()) { diff --git a/infini_train/src/kernels/cuda/linear.cu b/infini_train/src/kernels/cuda/linear.cu index 1b4c1819..7b87296b 100644 --- a/infini_train/src/kernels/cuda/linear.cu +++ b/infini_train/src/kernels/cuda/linear.cu @@ -136,22 +136,22 @@ std::shared_ptr LinearForward(const std::shared_ptr &input, cons } template -__global__ void ReduceColumnsKernel(const TIn *__restrict__ input, TOut *__restrict__ output, int num_rows, - int num_cols) { +__global__ void ReduceRowsKernel(const TIn *__restrict__ input, TOut *__restrict__ output, int64_t num_rows, + int64_t num_cols) { using BlockReduce = cub::BlockReduce; __shared__ typename BlockReduce::TempStorage temp_storage; - int row = blockIdx.x; + const int64_t col = blockIdx.x; float sum = 0.0f; - for (int col = threadIdx.x; col < num_cols; col += blockDim.x) { + for (int64_t row = threadIdx.x; row < num_rows; row += blockDim.x) { sum += common::cuda::Cast(input[row * num_cols + col]); } float reduced = BlockReduce(temp_storage).Sum(sum); if (threadIdx.x == 0) { - output[row] = reduced; + output[col] = common::cuda::Cast(reduced); } } @@ -289,7 +289,8 @@ std::shared_ptr LinearBackwardWeight(const std::shared_ptr &inpu std::shared_ptr LinearBackwardBias(const std::shared_ptr &grad_output, int64_t out_features) { const auto &dims = grad_output->Dims(); CHECK_GE(dims.size(), 2); - const int64_t bs = std::accumulate(dims.rbegin() + 1, dims.rend(), 1, std::multiplies{}); + CHECK_EQ(dims.back(), out_features); + const int64_t bs = std::accumulate(dims.rbegin() + 1, dims.rend(), int64_t{1}, std::multiplies{}); auto compute_dtype = grad_output->Dtype(); // FIXME(cx): output dtype promotion is a temporary hack; revisit when autograd/autocast is fixed. @@ -307,15 +308,15 @@ std::shared_ptr LinearBackwardBias(const std::shared_ptr &grad_o constexpr int BLOCK_SIZE = 256; switch (compute_dtype) { DISPATCH_CASE(WRAP({ - ReduceColumnsKernel<<>>( + ReduceRowsKernel<<>>( static_cast(grad_output->DataPtr()), - static_cast(grad_bias->DataPtr()), out_features, bs); + static_cast(grad_bias->DataPtr()), bs, out_features); }), DataType::kFLOAT32) DISPATCH_CASE(WRAP({ - ReduceColumnsKernel<<>>( + ReduceRowsKernel<<>>( static_cast(grad_output->DataPtr()), - static_cast(grad_bias->DataPtr()), out_features, bs); + static_cast(grad_bias->DataPtr()), bs, out_features); }), DataType::kBFLOAT16) } diff --git a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc index 57460ba0..19361e96 100644 --- a/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc +++ b/infini_train/src/nn/parallel/ddp/distributed_data_parallel.cc @@ -28,23 +28,30 @@ DistributedDataParallel::DistributedDataParallel(std::shared_ptr mod if (ddp_config_.zero_stage == 3) { LOG(FATAL) << "DistributedDataParallel: ZeRO-3 is not implemented yet."; } + CHECK_NOTNULL(ddp_pg_); + const auto expected_backend = ddp_pg_->backend(); + const int expected_device_index = global::GetDeviceIndex(rank.thread_rank()); + const auto validate_device = [expected_backend, expected_device_index](Device device, const char *kind) { + CHECK_EQ(static_cast(device.type()), static_cast(expected_backend)) + << "DistributedDataParallel " << kind << " backend must match the process group backend"; + CHECK_EQ(device.index(), expected_device_index) + << "DistributedDataParallel " << kind << " must use the device assigned to this rank"; + }; + for (auto ¶m : module->Parameters()) { + auto device = param->GetDevice(); + validate_device(device, "parameter"); if (!param->requires_grad()) { continue; } - auto device = param->GetDevice(); - CHECK_EQ(device.index(), global::GetDeviceIndex(rank.thread_rank())) - << "All parameters must be on the same device as the module"; if (!ddp_config.gradient_bucketing_enabled && ddp_config.zero_stage < 1) { - auto hook = std::make_unique( - function::ReduceOpType::kAvg, ddp_pg_); + const auto reduce_op + = ddp_config.average_in_collective ? function::ReduceOpType::kAvg : function::ReduceOpType::kSum; + auto hook = std::make_unique(reduce_op, ddp_pg_); param->RegisterPostAccumulateGradHook(std::move(hook)); } } - for (auto &buffer : module->Buffers()) { - CHECK_EQ(buffer->GetDevice().index(), global::GetDeviceIndex(rank.thread_rank())) - << "All buffers must be on the same device as the module"; - } + for (auto &buffer : module->Buffers()) { validate_device(buffer->GetDevice(), "buffer"); } modules_[kModuleName] = std::move(module); if (ddp_config.zero_stage >= 1) { diff --git a/infini_train/src/nn/parallel/ddp/reducer.cc b/infini_train/src/nn/parallel/ddp/reducer.cc index a13a3937..80115b8e 100644 --- a/infini_train/src/nn/parallel/ddp/reducer.cc +++ b/infini_train/src/nn/parallel/ddp/reducer.cc @@ -401,7 +401,9 @@ void Reducer::FinalizeBucketDense(size_t bucket_index) { // FIXME(zbl): support custom hook later LOG(FATAL) << "Custom hook is not supported now"; } else { - bucket.work = ddp_pg->AllReduce(bucket.contents, function::ReduceOpType::kAvg, true); + const auto reduce_op + = ddp_config_.average_in_collective ? function::ReduceOpType::kAvg : function::ReduceOpType::kSum; + bucket.work = ddp_pg->AllReduce(bucket.contents, reduce_op, true); } } diff --git a/infini_train/src/tensor.cc b/infini_train/src/tensor.cc index 18ca3d22..4e61e221 100644 --- a/infini_train/src/tensor.cc +++ b/infini_train/src/tensor.cc @@ -157,6 +157,8 @@ Tensor Tensor::To(Device device) { // 1. D2H Tensor cpu_tensor = To(Device()); // 2. H2D + // FIXME: Use the destination device for the guard, runtime implementation, and stream + // when cross-backend copies are supported. core::DeviceGuard guard(buffer_device); auto *impl = core::GetDeviceGuardImpl(buffer_device.type()); impl->MemcpyAsync(new_tensor.DataPtr(), cpu_tensor.DataPtr(), SizeInBytes(), core::MemcpyKind::kH2D, diff --git a/tests/autograd/test_autograd_elementwise_backward.cc b/tests/autograd/test_autograd_elementwise_backward.cc index f7eb0d5f..8623fe97 100644 --- a/tests/autograd/test_autograd_elementwise_backward.cc +++ b/tests/autograd/test_autograd_elementwise_backward.cc @@ -4,6 +4,7 @@ #include "gtest/gtest.h" #include "infini_train/include/autograd/elementwise.h" +#include "infini_train/include/core/runtime/device_guard.h" #include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/tensor.h" @@ -11,6 +12,27 @@ using namespace infini_train; +namespace { + +void ExpectExpGradient(const std::shared_ptr &actual, const std::vector &input_values, + const std::vector &grad_values, Device expected_device) { + ASSERT_NE(actual, nullptr); + ASSERT_EQ(input_values.size(), grad_values.size()); + ASSERT_EQ(actual->NumElements(), input_values.size()); + EXPECT_EQ(actual->Dtype(), DataType::kFLOAT32); + EXPECT_EQ(actual->GetDevice(), expected_device); + + auto actual_cpu = actual->To(Device()); + core::GetDeviceGuardImpl(actual->GetDevice().type())->SynchronizeDevice(actual->GetDevice()); + const auto *actual_data = static_cast(actual_cpu.DataPtr()); + for (size_t idx = 0; idx < input_values.size(); ++idx) { + const float expected = grad_values[idx] * std::exp(input_values[idx]); + EXPECT_NEAR(actual_data[idx], expected, 1e-5f) << "Mismatch at index " << idx; + } +} + +} // namespace + class AutogradElementwiseBackwardTest : public infini_train::test::InfiniTrainTest {}; TEST_P(AutogradElementwiseBackwardTest, AddBackward) { @@ -23,7 +45,9 @@ TEST_P(AutogradElementwiseBackwardTest, AddBackward) { auto grad = std::make_shared(std::vector{2, 3}, DataType::kFLOAT32, GetDevice(), true); grad->Fill(1.0f); auto grad_inputs = add_fn->Backward({grad}); - EXPECT_EQ(grad_inputs.size(), 2); + ASSERT_EQ(grad_inputs.size(), 2); + EXPECT_NE(grad_inputs[0].get(), grad_inputs[1].get()); + EXPECT_NE(grad_inputs[0]->DataPtr(), grad_inputs[1]->DataPtr()); } TEST_P(AutogradElementwiseBackwardTest, SubBackward) { @@ -126,14 +150,33 @@ TEST_P(AutogradElementwiseBackwardTest, TanhBackward) { } TEST_P(AutogradElementwiseBackwardTest, ExpBackward) { - auto a = std::make_shared(std::vector{2, 3}, DataType::kFLOAT32, GetDevice(), true); - a->Fill(1.0f); + const std::vector input_values = {-1.0f, -0.5f, 0.0f, 0.5f, 1.0f, 2.0f}; + const std::vector grad_values = {0.25f, -0.5f, 1.0f, 1.5f, -2.0f, 0.125f}; + auto a = std::make_shared(input_values.data(), std::vector{2, 3}, DataType::kFLOAT32, GetDevice()); + a->RequiresGrad(); auto exp_fn = std::make_shared(); auto result = exp_fn->Apply({a}); - auto grad = std::make_shared(std::vector{2, 3}, DataType::kFLOAT32, GetDevice(), true); - grad->Fill(1.0f); + ASSERT_EQ(result.size(), 1); + auto grad + = std::make_shared(grad_values.data(), std::vector{2, 3}, DataType::kFLOAT32, GetDevice()); auto grad_inputs = exp_fn->Backward({grad}); - EXPECT_EQ(grad_inputs.size(), 1); + ASSERT_EQ(grad_inputs.size(), 1); + ExpectExpGradient(grad_inputs[0], input_values, grad_values, GetDevice()); +} + +TEST_P(AutogradElementwiseBackwardTest, ExpBackwardAccumulatesIntoLeaf) { + const std::vector input_values = {-1.0f, -0.5f, 0.0f, 0.5f, 1.0f, 2.0f}; + const std::vector grad_values = {0.25f, -0.5f, 1.0f, 1.5f, -2.0f, 0.125f}; + auto input + = std::make_shared(input_values.data(), std::vector{2, 3}, DataType::kFLOAT32, GetDevice()); + input->RequiresGrad(); + auto output = input->Exp(); + auto grad + = std::make_shared(grad_values.data(), std::vector{2, 3}, DataType::kFLOAT32, GetDevice()); + + output->Backward(grad); + + ExpectExpGradient(input->grad(), input_values, grad_values, GetDevice()); } TEST_P(AutogradElementwiseBackwardTest, LogBackward) { diff --git a/tests/autograd/test_autograd_linear_backward.cc b/tests/autograd/test_autograd_linear_backward.cc index ba0f6fe1..9ce88eee 100644 --- a/tests/autograd/test_autograd_linear_backward.cc +++ b/tests/autograd/test_autograd_linear_backward.cc @@ -3,6 +3,7 @@ #include "gtest/gtest.h" #include "infini_train/include/autograd/linear.h" +#include "infini_train/include/core/runtime/device_guard.h" #include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/tensor.h" @@ -21,10 +22,17 @@ TEST_P(AutogradLinearBackwardTest, LinearBackward) { bias->Fill(0.0f); auto linear_fn = std::make_shared(); auto result = linear_fn->Apply({input, weight, bias}); - auto grad = std::make_shared(std::vector{2, 4}, DataType::kFLOAT32, GetDevice(), true); - grad->Fill(1.0f); + const float grad_values[] = {1.0f, 2.0f, 3.0f, 4.0f, 5.0f, 6.0f, 7.0f, 8.0f}; + auto grad = std::make_shared(grad_values, std::vector{2, 4}, DataType::kFLOAT32, GetDevice()); auto grad_inputs = linear_fn->Backward({grad}); - EXPECT_EQ(grad_inputs.size(), 3); + ASSERT_EQ(grad_inputs.size(), 3); + ASSERT_NE(grad_inputs[2], nullptr); + + auto bias_grad_cpu = grad_inputs[2]->To(Device()); + core::GetDeviceGuardImpl(GetDevice().type())->SynchronizeDevice(GetDevice()); + const auto *bias_grad = static_cast(bias_grad_cpu.DataPtr()); + const float expected_bias_grad[] = {6.0f, 8.0f, 10.0f, 12.0f}; + for (int idx = 0; idx < 4; ++idx) { EXPECT_FLOAT_EQ(bias_grad[idx], expected_bias_grad[idx]); } } TEST_P(AutogradLinearBackwardTest, LinearBackwardNoBias) { diff --git a/tests/optimizer/CMakeLists.txt b/tests/optimizer/CMakeLists.txt index bce88694..c0bfbd50 100644 --- a/tests/optimizer/CMakeLists.txt +++ b/tests/optimizer/CMakeLists.txt @@ -7,17 +7,3 @@ file(GLOB OPTIMIZER_SOURCES ${CMAKE_CURRENT_SOURCE_DIR}/test_*.cc) infini_train_add_test_suite(test_optimizer SOURCES ${OPTIMIZER_SOURCES} ) - -add_test( - NAME OptimizerParameterNamesTest.DistributedOptimizerPropagatesNamesToShardOptimizer - COMMAND ${CMAKE_COMMAND} -E env - "PROC_WORLD_SIZE=2" - $ - --gtest_filter=CUDA/OptimizerParameterNamesTest.DistributedOptimizerPropagatesNamesToShardOptimizer/* -) -set_tests_properties( - OptimizerParameterNamesTest.DistributedOptimizerPropagatesNamesToShardOptimizer - PROPERTIES - LABELS "cuda;distributed" - TIMEOUT 30 -) diff --git a/tests/optimizer/test_optimizer_parameter_names.cc b/tests/optimizer/test_optimizer_parameter_names.cc index 31943454..3b50b585 100644 --- a/tests/optimizer/test_optimizer_parameter_names.cc +++ b/tests/optimizer/test_optimizer_parameter_names.cc @@ -3,14 +3,6 @@ #include "gtest/gtest.h" -#include "infini_train/include/nn/modules/linear.h" -#include "infini_train/include/nn/parallel/ddp/distributed_data_parallel.h" -#include "infini_train/include/nn/parallel/ddp/distributed_data_parallel_config.h" -#include "infini_train/include/nn/parallel/ddp/distributed_optimizer.h" -#include "infini_train/include/nn/parallel/global.h" -#include "infini_train/include/nn/parallel/process_group.h" -#include "infini_train/include/nn/parallel/rank.h" -#include "infini_train/include/nn/parallel/utils.h" #include "infini_train/include/optimizer.h" #include "infini_train/include/tensor.h" @@ -52,38 +44,6 @@ TEST_P(OptimizerParameterNamesTest, ConstructorMatchesNamesToOptimizerParameterO EXPECT_TRUE(state.contains("adam.v.first")); } -TEST_P(OptimizerParameterNamesTest, DistributedOptimizerPropagatesNamesToShardOptimizer) { - ONLY_CUDA(); - REQUIRE_MIN_DEVICES(2); - if (nn::parallel::global::GetDataParallelSize() != 2) { - GTEST_SKIP() << "requires PROC_WORLD_SIZE=2"; - } - - const nn::parallel::Rank rank(/*process_rank=*/0, /*thread_rank=*/0, /*process_size=*/1, /*thread_size=*/2); - auto *pg_factory = nn::parallel::ProcessGroupFactory::Instance(Device::DeviceType::kCUDA); - pg_factory->GetOrCreate(nn::parallel::GetDataParallelProcessGroupName(rank.GlobalRank()), - nn::parallel::GetDataParallelGroupRanks(rank.GlobalRank())); - - auto model = std::make_shared(4, 4, /*bias=*/false, GetDevice()); - const auto named_parameters = model->NamedParameters(); - - nn::parallel::DistributedDataParallelConfig ddp_config; - ddp_config.zero_stage = 1; - ddp_config.overlap_grad_reduce = false; - ddp_config.overlap_param_gather = false; - auto ddp_model = std::make_shared(model, rank, ddp_config); - - nn::parallel::DistributedOptimizer optimizer(optimizers::Adam::CreateNamed(0.001), named_parameters, - std::vector>{ddp_model}, - /*ddp_world_size=*/2, /*ddp_rank=*/0); - const auto state = optimizer.StateDict(); - - EXPECT_TRUE(state.contains("adam.m.weight")); - EXPECT_TRUE(state.contains("adam.v.weight")); - EXPECT_FALSE(state.contains("adam.m.0")); - EXPECT_FALSE(state.contains("adam.v.0")); -} - TEST_P(OptimizerParameterNamesTest, PreservesNumericKeysWhenNamesAreNotSet) { auto parameter = std::make_shared(std::vector{2, 2}, DataType::kFLOAT32, GetDevice()); auto adam = std::make_shared(std::vector>{parameter}, 0.001);