diff --git a/infini_train/src/kernels/cuda/accumulate_grad.cu b/infini_train/src/kernels/cuda/accumulate_grad.cu index 93409a7e..91229265 100644 --- a/infini_train/src/kernels/cuda/accumulate_grad.cu +++ b/infini_train/src/kernels/cuda/accumulate_grad.cu @@ -39,22 +39,22 @@ void AccumulateGrad(const std::shared_ptr &gradient, float rate, const s "CUDA AccumulateGrad"); } -template -__global__ void AdamAccumulateGradKernel(const T *grad_data, T *param_data, size_t num_elements, T *m_data, T *v_data, - float learning_rate, float beta1, float beta2, float eps, +template +__global__ void AdamAccumulateGradKernel(const GradT *grad_data, ParamT *param_data, size_t num_elements, float *m_data, + float *v_data, float learning_rate, float beta1, float beta2, float eps, const float bias_correction_m, const float bias_correction_v) { size_t idx = blockIdx.x * blockDim.x + threadIdx.x; if (idx < num_elements) { - m_data[idx] = common::cuda::Fma(common::cuda::Cast(beta1), m_data[idx], - common::cuda::Cast(1 - beta1) * grad_data[idx]); - v_data[idx] = common::cuda::Fma(common::cuda::Cast(beta2), v_data[idx], - common::cuda::Cast(1 - beta2) * grad_data[idx] * grad_data[idx]); + const float grad = common::cuda::Cast(grad_data[idx]); + m_data[idx] = fmaf(beta1, m_data[idx], (1.0f - beta1) * grad); + v_data[idx] = fmaf(beta2, v_data[idx], (1.0f - beta2) * grad * grad); - const float m_hat = common::cuda::Cast(m_data[idx]) / bias_correction_m; - const float v_hat = common::cuda::Cast(v_data[idx]) / bias_correction_v; + const float m_hat = m_data[idx] / bias_correction_m; + const float v_hat = v_data[idx] / bias_correction_v; - param_data[idx] = common::cuda::Sub( - param_data[idx], common::cuda::Cast(learning_rate * m_hat * __frcp_rn(__fsqrt_rn(v_hat) + eps))); + const float param = common::cuda::Cast(param_data[idx]); + param_data[idx] + = common::cuda::Cast(param - learning_rate * m_hat * __frcp_rn(__fsqrt_rn(v_hat) + eps)); } } @@ -70,19 +70,26 @@ void AdamAccumulateGrad(const std::shared_ptr &grad, const std::shared_p int num_blocks = (num_elements + threads_per_block - 1) / threads_per_block; auto device = grad->GetDevice(); + CHECK_EQ(static_cast(m->Dtype()), static_cast(DataType::kFLOAT32)); + CHECK_EQ(static_cast(v->Dtype()), static_cast(DataType::kFLOAT32)); const auto &cuda_stream = dynamic_cast( infini_train::core::GetDeviceGuardImpl(device.type())->GetStream(device)) ->cuda_stream(); core::cuda::DispatchCudaFunc( grad->Dtype(), - [=]() { - AdamAccumulateGradKernel<<>>( - static_cast(grad->DataPtr()), static_cast(param->DataPtr()), num_elements, - static_cast(m->DataPtr()), static_cast(v->DataPtr()), learning_rate, beta1, beta2, eps, - bias_correction_m, bias_correction_v); + [=]() { + core::cuda::DispatchCudaFunc( + param->Dtype(), + [=]() { + AdamAccumulateGradKernel<<>>( + static_cast(grad->DataPtr()), static_cast(param->DataPtr()), + num_elements, static_cast(m->DataPtr()), static_cast(v->DataPtr()), + learning_rate, beta1, beta2, eps, bias_correction_m, bias_correction_v); + }, + "CUDA AdamAccumulateGrad parameter"); }, - "CUDA AdamAccumulateGrad"); + "CUDA AdamAccumulateGrad gradient"); } } // namespace infini_train::kernels::cuda diff --git a/infini_train/src/optimizer.cc b/infini_train/src/optimizer.cc index 39b999c7..3437fced 100644 --- a/infini_train/src/optimizer.cc +++ b/infini_train/src/optimizer.cc @@ -82,8 +82,8 @@ Adam::Adam(const std::vector> ¶ms, float learning_ra : Optimizer(params, learning_rate), t_(0), beta1_(beta1), beta2_(beta2), eps_(eps) { for (const auto ¶m : params_) { - m_.emplace_back(std::make_shared(param->Dims(), param->Dtype(), param->GetDevice())); - v_.emplace_back(std::make_shared(param->Dims(), param->Dtype(), param->GetDevice())); + m_.emplace_back(std::make_shared(param->Dims(), DataType::kFLOAT32, param->GetDevice())); + v_.emplace_back(std::make_shared(param->Dims(), DataType::kFLOAT32, param->GetDevice())); m_.back()->Fill(0.0); v_.back()->Fill(0.0); } @@ -92,8 +92,8 @@ Adam::Adam(const std::vector> ¶ms, float learning_ra Adam::Adam(const NamedParameterList &named_params, float learning_rate, float beta1, float beta2, float eps) : Optimizer(named_params, learning_rate), t_(0), beta1_(beta1), beta2_(beta2), eps_(eps) { for (const auto &[name, param] : named_params) { - m_.emplace_back(std::make_shared(param->Dims(), param->Dtype(), param->GetDevice())); - v_.emplace_back(std::make_shared(param->Dims(), param->Dtype(), param->GetDevice())); + m_.emplace_back(std::make_shared(param->Dims(), DataType::kFLOAT32, param->GetDevice())); + v_.emplace_back(std::make_shared(param->Dims(), DataType::kFLOAT32, param->GetDevice())); m_.back()->Fill(0.0); v_.back()->Fill(0.0); } diff --git a/tests/optimizer/test_optimizer_creation.cc b/tests/optimizer/test_optimizer_creation.cc index eac6d0b8..158a2686 100644 --- a/tests/optimizer/test_optimizer_creation.cc +++ b/tests/optimizer/test_optimizer_creation.cc @@ -25,6 +25,23 @@ TEST_P(OptimizerCreationTest, AdamCreation) { EXPECT_NE(optimizer, nullptr); } +TEST_P(OptimizerCreationTest, AdamStateIsFP32ForBF16Parameters) { + auto param = std::make_shared(std::vector{2, 3}, DataType::kBFLOAT16, GetDevice()); + auto optimizer = std::make_shared(std::vector>{param}, 0.001); + const auto state = optimizer->StateDict(); + EXPECT_EQ(state.at("adam.m.0")->Dtype(), DataType::kFLOAT32); + EXPECT_EQ(state.at("adam.v.0")->Dtype(), DataType::kFLOAT32); +} + +TEST_P(OptimizerCreationTest, NamedAdamStateIsFP32ForBF16Parameters) { + auto param = std::make_shared(std::vector{2, 3}, DataType::kBFLOAT16, GetDevice()); + const NamedParameterList named_params{{"weight", param}}; + auto optimizer = std::make_shared(named_params, 0.001); + const auto state = optimizer->StateDict(); + EXPECT_EQ(state.at("adam.m.weight")->Dtype(), DataType::kFLOAT32); + EXPECT_EQ(state.at("adam.v.weight")->Dtype(), DataType::kFLOAT32); +} + TEST_P(OptimizerCreationTest, SGDMultiParams) { std::vector> params; for (int i = 0; i < 3; ++i) { diff --git a/tests/optimizer/test_optimizer_step.cc b/tests/optimizer/test_optimizer_step.cc index 66ef1be7..41c83a6f 100644 --- a/tests/optimizer/test_optimizer_step.cc +++ b/tests/optimizer/test_optimizer_step.cc @@ -30,6 +30,26 @@ TEST_P(OptimizerStepTest, AdamStep) { optimizer->Step(); } +TEST_P(OptimizerStepTest, AdamUpdatesBF16ParameterWithFP32State) { + if (GetDevice().type() != Device::DeviceType::kCUDA) { + GTEST_SKIP() << "BF16 Adam update is CUDA-only"; + } + auto param = std::make_shared(std::vector{2, 3}, DataType::kBFLOAT16, GetDevice()); + param->set_requires_grad(true); + param->Fill(1.0f); + auto grad = std::make_shared(param->Dims(), DataType::kBFLOAT16, GetDevice()); + grad->Fill(0.5f); + param->set_grad(grad); + + auto optimizer = std::make_shared(std::vector>{param}, 0.01); + optimizer->Step(); + + const auto updated = param->To(DataType::kFLOAT32).To(Device()); + const auto *data = static_cast(updated.DataPtr()); + EXPECT_LT(data[0], 1.0f); + EXPECT_EQ(optimizer->StateDict().at("adam.m.0")->Dtype(), DataType::kFLOAT32); +} + TEST_P(OptimizerStepTest, ZeroGrad) { auto param = std::make_shared(std::vector{2, 3}, DataType::kFLOAT32, GetDevice()); param->set_requires_grad(true);