Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
41 changes: 24 additions & 17 deletions infini_train/src/kernels/cuda/accumulate_grad.cu
Original file line number Diff line number Diff line change
Expand Up @@ -39,22 +39,22 @@ void AccumulateGrad(const std::shared_ptr<Tensor> &gradient, float rate, const s
"CUDA AccumulateGrad");
}

template <typename T>
__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 <typename GradT, typename ParamT>
__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<T>(beta1), m_data[idx],
common::cuda::Cast<T>(1 - beta1) * grad_data[idx]);
v_data[idx] = common::cuda::Fma(common::cuda::Cast<T>(beta2), v_data[idx],
common::cuda::Cast<T>(1 - beta2) * grad_data[idx] * grad_data[idx]);
const float grad = common::cuda::Cast<float>(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<float>(m_data[idx]) / bias_correction_m;
const float v_hat = common::cuda::Cast<float>(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<T>(learning_rate * m_hat * __frcp_rn(__fsqrt_rn(v_hat) + eps)));
const float param = common::cuda::Cast<float>(param_data[idx]);
param_data[idx]
= common::cuda::Cast<ParamT>(param - learning_rate * m_hat * __frcp_rn(__fsqrt_rn(v_hat) + eps));
}
}

Expand All @@ -70,19 +70,26 @@ void AdamAccumulateGrad(const std::shared_ptr<Tensor> &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<int>(m->Dtype()), static_cast<int>(DataType::kFLOAT32));
CHECK_EQ(static_cast<int>(v->Dtype()), static_cast<int>(DataType::kFLOAT32));
const auto &cuda_stream = dynamic_cast<infini_train::core::cuda::CudaStream *>(
infini_train::core::GetDeviceGuardImpl(device.type())->GetStream(device))
->cuda_stream();

core::cuda::DispatchCudaFunc<INFINI_ALL_FLOATING_TYPES>(
grad->Dtype(),
[=]<typename T>() {
AdamAccumulateGradKernel<<<num_blocks, threads_per_block, 0, cuda_stream>>>(
static_cast<const T *>(grad->DataPtr()), static_cast<T *>(param->DataPtr()), num_elements,
static_cast<T *>(m->DataPtr()), static_cast<T *>(v->DataPtr()), learning_rate, beta1, beta2, eps,
bias_correction_m, bias_correction_v);
[=]<typename GradT>() {
core::cuda::DispatchCudaFunc<INFINI_ALL_FLOATING_TYPES>(
param->Dtype(),
[=]<typename ParamT>() {
AdamAccumulateGradKernel<<<num_blocks, threads_per_block, 0, cuda_stream>>>(
static_cast<const GradT *>(grad->DataPtr()), static_cast<ParamT *>(param->DataPtr()),
num_elements, static_cast<float *>(m->DataPtr()), static_cast<float *>(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

Expand Down
8 changes: 4 additions & 4 deletions infini_train/src/optimizer.cc
Original file line number Diff line number Diff line change
Expand Up @@ -82,8 +82,8 @@ Adam::Adam(const std::vector<std::shared_ptr<Tensor>> &params, float learning_ra
: Optimizer(params, learning_rate), t_(0), beta1_(beta1), beta2_(beta2), eps_(eps) {

for (const auto &param : params_) {
m_.emplace_back(std::make_shared<Tensor>(param->Dims(), param->Dtype(), param->GetDevice()));
v_.emplace_back(std::make_shared<Tensor>(param->Dims(), param->Dtype(), param->GetDevice()));
m_.emplace_back(std::make_shared<Tensor>(param->Dims(), DataType::kFLOAT32, param->GetDevice()));
v_.emplace_back(std::make_shared<Tensor>(param->Dims(), DataType::kFLOAT32, param->GetDevice()));
m_.back()->Fill(0.0);
v_.back()->Fill(0.0);
}
Expand All @@ -92,8 +92,8 @@ Adam::Adam(const std::vector<std::shared_ptr<Tensor>> &params, 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<Tensor>(param->Dims(), param->Dtype(), param->GetDevice()));
v_.emplace_back(std::make_shared<Tensor>(param->Dims(), param->Dtype(), param->GetDevice()));
m_.emplace_back(std::make_shared<Tensor>(param->Dims(), DataType::kFLOAT32, param->GetDevice()));
v_.emplace_back(std::make_shared<Tensor>(param->Dims(), DataType::kFLOAT32, param->GetDevice()));
m_.back()->Fill(0.0);
v_.back()->Fill(0.0);
}
Expand Down
17 changes: 17 additions & 0 deletions tests/optimizer/test_optimizer_creation.cc
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,23 @@ TEST_P(OptimizerCreationTest, AdamCreation) {
EXPECT_NE(optimizer, nullptr);
}

TEST_P(OptimizerCreationTest, AdamStateIsFP32ForBF16Parameters) {
auto param = std::make_shared<Tensor>(std::vector<int64_t>{2, 3}, DataType::kBFLOAT16, GetDevice());
auto optimizer = std::make_shared<optimizers::Adam>(std::vector<std::shared_ptr<Tensor>>{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<Tensor>(std::vector<int64_t>{2, 3}, DataType::kBFLOAT16, GetDevice());
const NamedParameterList named_params{{"weight", param}};
auto optimizer = std::make_shared<optimizers::Adam>(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<std::shared_ptr<Tensor>> params;
for (int i = 0; i < 3; ++i) {
Expand Down
20 changes: 20 additions & 0 deletions tests/optimizer/test_optimizer_step.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<Tensor>(std::vector<int64_t>{2, 3}, DataType::kBFLOAT16, GetDevice());
param->set_requires_grad(true);
param->Fill(1.0f);
auto grad = std::make_shared<Tensor>(param->Dims(), DataType::kBFLOAT16, GetDevice());
grad->Fill(0.5f);
param->set_grad(grad);

auto optimizer = std::make_shared<optimizers::Adam>(std::vector<std::shared_ptr<Tensor>>{param}, 0.01);
optimizer->Step();

const auto updated = param->To(DataType::kFLOAT32).To(Device());
const auto *data = static_cast<const float *>(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<Tensor>(std::vector<int64_t>{2, 3}, DataType::kFLOAT32, GetDevice());
param->set_requires_grad(true);
Expand Down
Loading