diff --git a/example/llama3/checkpoint_loader.cc b/example/llama3/checkpoint_loader.cc index f3590af6..0dead4df 100644 --- a/example/llama3/checkpoint_loader.cc +++ b/example/llama3/checkpoint_loader.cc @@ -277,7 +277,8 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; - ReadMatrixRowShardFloat(ifs, static_cast(tensor->DataPtr()), + float *dst = static_cast(tensor->DataPtr()) + fc_pp * n_embd; + ReadMatrixRowShardFloat(ifs, dst, /*rows=*/fc_out, /*cols=*/n_embd, /*row_start=*/tp_rank * fc_pp, /*row_cnt=*/fc_pp); ++local_layer_index; @@ -293,7 +294,7 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath) if (owned_layers[i]) { auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index), - nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFc2LayerName, + nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)]; ReadMatrixRowShardFloat(ifs, static_cast(tensor->DataPtr()), /*rows=*/fc_out, /*cols=*/n_embd, diff --git a/example/mixtral/checkpoint_loader.cc b/example/mixtral/checkpoint_loader.cc index c6c8471b..c451ce3c 100644 --- a/example/mixtral/checkpoint_loader.cc +++ b/example/mixtral/checkpoint_loader.cc @@ -119,10 +119,10 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath, CHECK(ifs) << "Failed to read tensor " << name; }; - auto read_projection_into_packed_qkv = [&](const std::string &packed_qkv_name, int64_t row_offset, int64_t num_rows, - const std::string &projection_name) { - CHECK(state.contains(packed_qkv_name)) << "Model state_dict does not contain " << packed_qkv_name; - std::shared_ptr tensor = state.at(packed_qkv_name); + auto read_projection_into_packed_weight = [&](const std::string &packed_weight_name, int64_t row_offset, + int64_t num_rows, const std::string &projection_name) { + CHECK(state.contains(packed_weight_name)) << "Model state_dict does not contain " << packed_weight_name; + std::shared_ptr tensor = state.at(packed_weight_name); CHECK(tensor->Dtype() == infini_train::DataType::kFLOAT32) << "Only float32 tiny Mixtral LLMC files are supported: " << projection_name; CHECK_EQ(tensor->Dims().size(), 2); @@ -144,17 +144,21 @@ std::shared_ptr LoadFromLLMC(const std::string &filepath, const int64_t head_dim = config.n_embd / config.n_head; const int64_t q_rows = config.n_head * head_dim; const int64_t kv_rows = config.n_kv_head * head_dim; - read_projection_into_packed_qkv(c_attn_name, 0, q_rows, c_attn_name + ".q_proj"); - read_projection_into_packed_qkv(c_attn_name, q_rows, kv_rows, c_attn_name + ".k_proj"); - read_projection_into_packed_qkv(c_attn_name, q_rows + kv_rows, kv_rows, c_attn_name + ".v_proj"); + read_projection_into_packed_weight(c_attn_name, 0, q_rows, c_attn_name + ".q_proj"); + read_projection_into_packed_weight(c_attn_name, q_rows, kv_rows, c_attn_name + ".k_proj"); + read_projection_into_packed_weight(c_attn_name, q_rows + kv_rows, kv_rows, c_attn_name + ".v_proj"); read_tensor_by_state_key(prefix + ".attn.c_proj.weight"); read_tensor_by_state_key(prefix + ".ln_2.weight"); read_tensor_by_state_key(prefix + ".mlp.router.weight"); for (int64_t expert = 0; expert < moe_config.num_experts; ++expert) { const std::string expert_prefix = prefix + ".mlp.experts.expert_" + std::to_string(expert); - read_tensor_by_state_key(expert_prefix + ".c_fc2.weight"); // Mixtral w1/gate_proj - read_tensor_by_state_key(expert_prefix + ".c_fc.weight"); // Mixtral w3/up_proj - read_tensor_by_state_key(expert_prefix + ".c_proj.weight"); // Mixtral w2/down_proj + const std::string packed_fc1_name = expert_prefix + ".c_fc.weight"; + read_projection_into_packed_weight(packed_fc1_name, 0, moe_config.moe_ffn_hidden_size, + expert_prefix + ".c_fc2.weight"); // Mixtral w1/gate_proj + read_projection_into_packed_weight(packed_fc1_name, moe_config.moe_ffn_hidden_size, + moe_config.moe_ffn_hidden_size, + expert_prefix + ".c_fc.weight"); // Mixtral w3/up_proj + read_tensor_by_state_key(expert_prefix + ".c_proj.weight"); // Mixtral w2/down_proj } } read_tensor_by_state_key("transformer.ln_f.weight"); diff --git a/infini_train/include/autograd/activations.h b/infini_train/include/autograd/activations.h index a6397726..079db0fb 100644 --- a/infini_train/include/autograd/activations.h +++ b/infini_train/include/autograd/activations.h @@ -21,4 +21,16 @@ class Sigmoid : public Function { const std::vector> &output_tensors) override; std::vector> Backward(const std::vector> &grad_outputs) override; }; + +class SwiGLU : public Function { +public: + static constexpr char kType[] = "SwiGLUFunction"; + + SwiGLU() : 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; +}; } // namespace infini_train::autograd diff --git a/infini_train/include/nn/modules/activations.h b/infini_train/include/nn/modules/activations.h index deb02957..549d9375 100644 --- a/infini_train/include/nn/modules/activations.h +++ b/infini_train/include/nn/modules/activations.h @@ -30,6 +30,7 @@ class SwiGLU : public CloneableModule { static constexpr char kType[] = "SwiGLU"; SwiGLU() : CloneableModule(kType) {} + // The last input dimension is packed as [gate, up], matching Megatron-LM. std::vector> Forward(const std::vector> &x) override; }; } // namespace infini_train::nn diff --git a/infini_train/include/nn/modules/transformer/mlp.h b/infini_train/include/nn/modules/transformer/mlp.h index bb096b7c..ecf5672b 100644 --- a/infini_train/include/nn/modules/transformer/mlp.h +++ b/infini_train/include/nn/modules/transformer/mlp.h @@ -13,8 +13,7 @@ class MLP : public infini_train::nn::CloneableModule { static constexpr char kGeluLayerName[] = "gelu"; static constexpr char kCProjLayerName[] = "c_proj"; - static constexpr char kCFc2LayerName[] = "c_fc2"; - static constexpr char kSiluLayerName[] = "silu"; + static constexpr char kSwiGLULayerName[] = "swiglu"; explicit MLP(const TransformerConfig &config); diff --git a/infini_train/src/autograd/activations.cc b/infini_train/src/autograd/activations.cc index bb8b8e5e..6894788a 100644 --- a/infini_train/src/autograd/activations.cc +++ b/infini_train/src/autograd/activations.cc @@ -30,4 +30,30 @@ std::vector> Sigmoid::Backward(const std::vectorGetDevice().type(); return {Dispatcher::Instance().Call>({device, "SigmoidBackward"}, output, grad_output)}; } + +std::vector> SwiGLU::Forward(const std::vector> &input_tensors) { + CHECK_EQ(input_tensors.size(), 1); + const auto &input = input_tensors[0]; + CHECK_GT(input->Dims().size(), 0); + CHECK_EQ(input->Dims().back() % 2, 0) << "SwiGLU expects an even last dimension"; + + auto device = input->GetDevice().type(); + return {Dispatcher::Instance().Call>({device, "SwiGLUForward"}, input)}; +} + +void SwiGLU::SetupContext(const std::vector> &input_tensors, + const std::vector> &) { + ctx_.SaveForBackward({input_tensors[0]}); +} + +std::vector> SwiGLU::Backward(const std::vector> &grad_outputs) { + auto saved_tensors = ctx_.GetSavedTensors(); + CHECK_EQ(saved_tensors.size(), 1); + CHECK_EQ(grad_outputs.size(), 1); + const auto &input = saved_tensors[0]; + const auto &grad_output = grad_outputs[0]; + + auto device = input->GetDevice().type(); + return {Dispatcher::Instance().Call>({device, "SwiGLUBackward"}, input, grad_output)}; +} } // namespace infini_train::autograd diff --git a/infini_train/src/kernels/cpu/swiglu.cc b/infini_train/src/kernels/cpu/swiglu.cc new file mode 100644 index 00000000..c87500c0 --- /dev/null +++ b/infini_train/src/kernels/cpu/swiglu.cc @@ -0,0 +1,74 @@ +#include +#include + +#include "glog/logging.h" + +#include "infini_train/include/dispatcher.h" +#include "infini_train/include/tensor.h" + +namespace infini_train::kernels::cpu { +std::shared_ptr SwiGLUForward(const std::shared_ptr &input) { + CHECK(input->Dtype() == DataType::kFLOAT32); + CHECK(input->IsContiguous()); + auto output_dims = input->Dims(); + CHECK_GT(output_dims.size(), 0); + CHECK_EQ(output_dims.back() % 2, 0); + const int64_t hidden = output_dims.back() / 2; + CHECK_GT(hidden, 0); + output_dims.back() = hidden; + + auto output = std::make_shared(output_dims, input->Dtype(), input->GetDevice()); + const float *input_ptr = static_cast(input->DataPtr()); + float *output_ptr = static_cast(output->DataPtr()); + const int64_t rows = output->NumElements() / hidden; + for (int64_t row = 0; row < rows; ++row) { + const int64_t input_base = row * 2 * hidden; + const int64_t output_base = row * hidden; + for (int64_t col = 0; col < hidden; ++col) { + const float gate = input_ptr[input_base + col]; + const float up = input_ptr[input_base + hidden + col]; + output_ptr[output_base + col] = up * gate / (1.0f + std::exp(-gate)); + } + } + return output; +} + +std::shared_ptr SwiGLUBackward(const std::shared_ptr &input, + const std::shared_ptr &grad_output) { + CHECK(input->Dtype() == DataType::kFLOAT32); + CHECK(grad_output->Dtype() == input->Dtype()); + CHECK(input->IsContiguous()); + CHECK(grad_output->IsContiguous()); + CHECK_GT(input->Dims().size(), 0); + const int64_t hidden = input->Dims().back() / 2; + CHECK_GT(hidden, 0); + CHECK_EQ(grad_output->NumElements() * 2, input->NumElements()); + + auto grad_input = std::make_shared(input->Dims(), input->Dtype(), input->GetDevice()); + const float *input_ptr = static_cast(input->DataPtr()); + const float *grad_output_ptr = static_cast(grad_output->DataPtr()); + float *grad_input_ptr = static_cast(grad_input->DataPtr()); + const int64_t rows = grad_output->NumElements() / hidden; + for (int64_t row = 0; row < rows; ++row) { + const int64_t input_base = row * 2 * hidden; + const int64_t output_base = row * hidden; + for (int64_t col = 0; col < hidden; ++col) { + const float gate = input_ptr[input_base + col]; + const float up = input_ptr[input_base + hidden + col]; + const float grad = grad_output_ptr[output_base + col]; + const float sigmoid = 1.0f / (1.0f + std::exp(-gate)); + grad_input_ptr[input_base + col] = grad * up * sigmoid * (1.0f + gate * (1.0f - sigmoid)); + grad_input_ptr[input_base + hidden + col] = grad * gate * sigmoid; + } + } + return grad_input; +} +} // namespace infini_train::kernels::cpu + +#define REGISTER_CPU_SWIGLU_KERNEL(kernel_name) \ + REGISTER_KERNEL(infini_train::Device::DeviceType::kCPU, kernel_name, infini_train::kernels::cpu::kernel_name) + +REGISTER_CPU_SWIGLU_KERNEL(SwiGLUForward) +REGISTER_CPU_SWIGLU_KERNEL(SwiGLUBackward) + +#undef REGISTER_CPU_SWIGLU_KERNEL diff --git a/infini_train/src/kernels/cuda/swiglu.cu b/infini_train/src/kernels/cuda/swiglu.cu new file mode 100644 index 00000000..e5fd8ab3 --- /dev/null +++ b/infini_train/src/kernels/cuda/swiglu.cu @@ -0,0 +1,147 @@ +#include +#include + +#include "infini_train/include/common/common.h" +#include "infini_train/include/common/cuda/kernel_helper.cuh" +#include "infini_train/include/core/runtime/device_guard.h" +#include "infini_train/include/datatype.h" +#include "infini_train/include/dispatcher.h" +#include "infini_train/include/tensor.h" + +#include "infini_train/src/core/runtime/cuda/cuda_runtime_common.h" + +namespace infini_train::kernels::cuda { +namespace { +using namespace infini_train::common::cuda; + +template +__global__ void SwiGLUForwardKernel(T *__restrict__ output, const T *__restrict__ input, int64_t hidden, + size_t num_elements) { + const size_t grid_stride = static_cast(gridDim.x) * blockDim.x; + for (size_t idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; idx < num_elements; + idx += grid_stride) { + const size_t row = idx / hidden; + const size_t col = idx % hidden; + const size_t input_base = row * 2 * hidden; + const T gate = input[input_base + col]; + const T up = input[input_base + hidden + col]; + output[idx] = Mul(up, Mul(gate, Sigmoid(gate))); + } +} + +template +__global__ void SwiGLUBackwardKernel(T *__restrict__ grad_input, const InputT *__restrict__ input, + const GradT *__restrict__ grad_output, int64_t hidden, size_t num_elements) { + const size_t grid_stride = static_cast(gridDim.x) * blockDim.x; + for (size_t idx = static_cast(blockIdx.x) * blockDim.x + threadIdx.x; idx < num_elements; + idx += grid_stride) { + const size_t row = idx / hidden; + const size_t col = idx % hidden; + const size_t input_base = row * 2 * hidden; + const T gate = Cast(input[input_base + col]); + const T up = Cast(input[input_base + hidden + col]); + const T grad = Cast(grad_output[idx]); + const T sigmoid = Sigmoid(gate); + grad_input[input_base + col] = Mul(grad, Mul(up, Mul(sigmoid, Add(T(1), Mul(gate, Sub(T(1), sigmoid)))))); + grad_input[input_base + hidden + col] = Mul(grad, Mul(gate, sigmoid)); + } +} + +inline size_t ChooseBlockSize(size_t num_elements) { + if (num_elements < 1024) { + return 64; + } + if (num_elements < 65536) { + return 128; + } + if (num_elements < 1048576) { + return 256; + } + return 512; +} +} // namespace + +std::shared_ptr SwiGLUForward(const std::shared_ptr &input) { + CHECK(input->IsContiguous()); + auto output_dims = input->Dims(); + CHECK_GT(output_dims.size(), 0); + CHECK_EQ(output_dims.back() % 2, 0); + const int64_t hidden = output_dims.back() / 2; + CHECK_GT(hidden, 0); + output_dims.back() = hidden; + + auto output = std::make_shared(output_dims, input->Dtype(), input->GetDevice()); + auto device = output->GetDevice(); + const auto &stream = dynamic_cast( + infini_train::core::GetDeviceGuardImpl(device.type())->GetStream(device)) + ->cuda_stream(); + const size_t num_elements = output->NumElements(); + const dim3 block(ChooseBlockSize(num_elements)); + const dim3 grid(std::min(CEIL_DIV(num_elements, block.x), static_cast(65535))); + + switch (input->Dtype()) { + case DataType::kFLOAT32: + SwiGLUForwardKernel<<>>(static_cast(output->DataPtr()), + static_cast(input->DataPtr()), hidden, + num_elements); + break; + case DataType::kBFLOAT16: + SwiGLUForwardKernel<<>>(static_cast(output->DataPtr()), + static_cast(input->DataPtr()), hidden, + num_elements); + break; + default: + LOG_LOC(FATAL, "CUDA SwiGLUForward: unsupported data type"); + } + return output; +} + +std::shared_ptr SwiGLUBackward(const std::shared_ptr &input, + const std::shared_ptr &grad_output) { + CHECK(input->IsContiguous()); + CHECK(grad_output->IsContiguous()); + CHECK_GT(input->Dims().size(), 0); + const int64_t hidden = input->Dims().back() / 2; + CHECK_GT(hidden, 0); + CHECK_EQ(grad_output->NumElements() * 2, input->NumElements()); + + const DataType output_dtype = PromoteDataTypes(input->Dtype(), grad_output->Dtype()); + auto grad_input = std::make_shared(input->Dims(), output_dtype, input->GetDevice()); + auto device = input->GetDevice(); + const auto &stream = dynamic_cast( + infini_train::core::GetDeviceGuardImpl(device.type())->GetStream(device)) + ->cuda_stream(); + const size_t num_elements = grad_output->NumElements(); + const dim3 block(ChooseBlockSize(num_elements)); + const dim3 grid(std::min(CEIL_DIV(num_elements, block.x), static_cast(65535))); + + if (input->Dtype() == DataType::kFLOAT32 && grad_output->Dtype() == DataType::kFLOAT32) { + SwiGLUBackwardKernel<<>>( + static_cast(grad_input->DataPtr()), static_cast(input->DataPtr()), + static_cast(grad_output->DataPtr()), hidden, num_elements); + } else if (input->Dtype() == DataType::kBFLOAT16 && grad_output->Dtype() == DataType::kBFLOAT16) { + SwiGLUBackwardKernel<<>>( + static_cast(grad_input->DataPtr()), static_cast(input->DataPtr()), + static_cast(grad_output->DataPtr()), hidden, num_elements); + } else if (input->Dtype() == DataType::kBFLOAT16 && grad_output->Dtype() == DataType::kFLOAT32) { + SwiGLUBackwardKernel<<>>( + static_cast(grad_input->DataPtr()), static_cast(input->DataPtr()), + static_cast(grad_output->DataPtr()), hidden, num_elements); + } else if (input->Dtype() == DataType::kFLOAT32 && grad_output->Dtype() == DataType::kBFLOAT16) { + SwiGLUBackwardKernel<<>>( + static_cast(grad_input->DataPtr()), static_cast(input->DataPtr()), + static_cast(grad_output->DataPtr()), hidden, num_elements); + } else { + LOG_LOC(FATAL, "CUDA SwiGLUBackward: unsupported data type combination"); + } + return grad_input; +} +} // namespace infini_train::kernels::cuda + +#define REGISTER_CUDA_SWIGLU_KERNEL(kernel_name) \ + REGISTER_KERNEL(infini_train::Device::DeviceType::kCUDA, kernel_name, infini_train::kernels::cuda::kernel_name) + +REGISTER_CUDA_SWIGLU_KERNEL(SwiGLUForward) +REGISTER_CUDA_SWIGLU_KERNEL(SwiGLUBackward) + +#undef REGISTER_CUDA_SWIGLU_KERNEL diff --git a/infini_train/src/nn/modules/activations.cc b/infini_train/src/nn/modules/activations.cc index d1bbc9da..37a57805 100644 --- a/infini_train/src/nn/modules/activations.cc +++ b/infini_train/src/nn/modules/activations.cc @@ -19,6 +19,6 @@ std::vector> NewGELU::Forward(const std::vector> SwiGLU::Forward(const std::vector> &x) { - return {x[0] * function::Sigmoid(x[0])}; + return std::make_shared()->Apply(x); } } // namespace infini_train::nn diff --git a/infini_train/src/nn/modules/transformer/mlp.cc b/infini_train/src/nn/modules/transformer/mlp.cc index 115d77e6..3cf94f91 100644 --- a/infini_train/src/nn/modules/transformer/mlp.cc +++ b/infini_train/src/nn/modules/transformer/mlp.cc @@ -48,29 +48,19 @@ MLP::MLP(const TransformerConfig &config) : CloneableModule(kType) { // c_fc: ColumnParallel (input full, output parallel) modules_[kCFcLayerName] = std::make_shared( - /*in_features=*/config.n_embd, /*out_features=*/ffn_hidden, + /*in_features=*/config.n_embd, + /*out_features=*/config.activation_type == MLPType::kSwiGLU ? 2 * ffn_hidden : ffn_hidden, /*bias=*/config.add_bias_linear, /*gather_output=*/false, /*input_is_parallel=*/false, /*skip_bias_add=*/false, /*sequence_parallel=*/parallel::global::GetSequenceParallelEnabled()); - // For SwiGLU, add second projection - if (config.activation_type == MLPType::kSwiGLU) { - modules_[kCFc2LayerName] = std::make_shared( - /*in_features=*/config.n_embd, /*out_features=*/ffn_hidden, - /*bias=*/config.add_bias_linear, - /*gather_output=*/false, - /*input_is_parallel=*/false, - /*skip_bias_add=*/false, - /*sequence_parallel=*/nn::parallel::global::GetSequenceParallelEnabled()); - } - // Activation: check for GELU or SwiGLU if (config.activation_type == MLPType::kGELU) { modules_[kGeluLayerName] = std::make_shared(); } else if (config.activation_type == MLPType::kSwiGLU) { - modules_[kSiluLayerName] = std::make_shared(); + modules_[kSwiGLULayerName] = std::make_shared(); } // c_proj: RowParallel (input parallel, output full) @@ -85,31 +75,22 @@ MLP::MLP(const TransformerConfig &config) : CloneableModule(kType) { std::vector> MLP::Forward(const std::vector> &x) { - bool is_swiglu = modules_.contains(kCFc2LayerName) && modules_.contains(kSiluLayerName); - - if (is_swiglu) { - // SwiGLU forward pass - // (B, T, C) -> ColumnParallelLinear(C, hidden_dim) -> (B, T, hidden_dim) - auto x1 = (*modules_[kCFcLayerName])(x)[0]; - // (B, T, C) -> ColumnParallelLinear(C, hidden_dim) -> (B, T, hidden_dim) - auto x2 = (*modules_[kCFc2LayerName])(x)[0]; - // (B, T, hidden_dim) -> SiLU -> (B, T, hidden_dim) - x2 = (*modules_[kSiluLayerName])({x2})[0]; - // (B, T, hidden_dim) -> element-wise mul -> (B, T, hidden_dim) - auto x3 = x1 * x2; - // (B, T, hidden_dim) -> RowParallelLinear(hidden_dim, C) -> (B, T, C) - auto x4 = (*modules_[kCProjLayerName])({x3}); - return x4; - } else { - // GELU forward pass (standard) - // (B, T, C) -> ColumnParallelLinear(C, 4*C) -> (B, T, 4*C_local) - auto x1 = (*modules_[kCFcLayerName])(x); - // (B, T, 4*C_local) -> GELU -> (B, T, 4*C_local) - auto x2 = (*modules_[kGeluLayerName])(x1); - // (B, T, 4*C_local) -> RowParallelLinear(4*C, C) -> (B, T, C) - auto x3 = (*modules_[kCProjLayerName])(x2); - return x3; + if (modules_.contains(kSwiGLULayerName)) { + // (B, T, C) -> ColumnParallelLinear(C, 2*H) -> (B, T, 2*H_local) + auto packed = (*modules_[kCFcLayerName])(x)[0]; + // (B, T, 2*H_local) [gate, up] -> SwiGLU -> (B, T, H_local) + auto activated = (*modules_[kSwiGLULayerName])({packed}); + // (B, T, H_local) -> RowParallelLinear(H, C) -> (B, T, C) + return (*modules_[kCProjLayerName])(activated); } + + // GELU forward pass (standard) + // (B, T, C) -> ColumnParallelLinear(C, 4*C) -> (B, T, 4*C_local) + auto x1 = (*modules_[kCFcLayerName])(x); + // (B, T, 4*C_local) -> GELU -> (B, T, 4*C_local) + auto x2 = (*modules_[kGeluLayerName])(x1); + // (B, T, 4*C_local) -> RowParallelLinear(4*C, C) -> (B, T, C) + return (*modules_[kCProjLayerName])(x2); } } // namespace infini_train::nn diff --git a/tests/autograd/test_autograd_elementwise_backward.cc b/tests/autograd/test_autograd_elementwise_backward.cc index f7eb0d5f..2e9d9a73 100644 --- a/tests/autograd/test_autograd_elementwise_backward.cc +++ b/tests/autograd/test_autograd_elementwise_backward.cc @@ -3,6 +3,7 @@ #include "gtest/gtest.h" +#include "infini_train/include/autograd/activations.h" #include "infini_train/include/autograd/elementwise.h" #include "infini_train/include/nn/parallel/global.h" #include "infini_train/include/tensor.h" @@ -13,6 +14,70 @@ using namespace infini_train; class AutogradElementwiseBackwardTest : public infini_train::test::InfiniTrainTest {}; +TEST_P(AutogradElementwiseBackwardTest, SwiGLUForwardBackward) { + const std::vector input_dims{2, 6}; + const std::vector input_values{-1.0f, 0.5f, 2.0f, -2.0f, 0.0f, 1.5f, 0.25f, -3.0f, 1.0f, 0.5f, -1.0f, 2.0f}; + const std::vector grad_values{1.0f, -0.5f, 2.0f, -1.5f, 0.25f, 0.75f}; + auto input = std::make_shared(input_values.data(), input_dims, DataType::kFLOAT32, GetDevice()); + auto grad_output + = std::make_shared(grad_values.data(), std::vector{2, 3}, DataType::kFLOAT32, GetDevice()); + + std::vector expected_output(6); + std::vector expected_grad(12); + for (int64_t row = 0; row < 2; ++row) { + for (int64_t col = 0; col < 3; ++col) { + const int64_t packed_base = row * 6; + const int64_t output_idx = row * 3 + col; + const float gate = input_values[packed_base + col]; + const float up = input_values[packed_base + 3 + col]; + const float grad = grad_values[output_idx]; + const float sigmoid = 1.0f / (1.0f + std::exp(-gate)); + expected_output[output_idx] = up * gate * sigmoid; + expected_grad[packed_base + col] = grad * up * sigmoid * (1.0f + gate * (1.0f - sigmoid)); + expected_grad[packed_base + 3 + col] = grad * gate * sigmoid; + } + } + + auto swiglu_fn = std::make_shared(); + auto result = swiglu_fn->Apply({input}); + ASSERT_EQ(result.size(), 1); + EXPECT_EQ(result[0]->Dims(), (std::vector{2, 3})); + test::ExpectTensorNear(result[0], expected_output, 1e-5f); + + auto grad_inputs = swiglu_fn->Backward({grad_output}); + ASSERT_EQ(grad_inputs.size(), 1); + EXPECT_EQ(grad_inputs[0]->Dims(), input_dims); + test::ExpectTensorNear(grad_inputs[0], expected_grad, 1e-5f); +} + +TEST_P(AutogradElementwiseBackwardTest, SwiGLUAutocastBackward) { + ONLY_CUDA(); + const std::vector input_dims{1, 4}; + const std::vector input_values{0.5f, -1.0f, 1.0f, -0.5f}; + const std::vector grad_values{2.0f, -0.25f}; + auto input_fp32 = std::make_shared(input_values.data(), input_dims, DataType::kFLOAT32, GetDevice()); + auto input = std::make_shared(input_fp32->To(DataType::kBFLOAT16)); + auto grad_output + = std::make_shared(grad_values.data(), std::vector{1, 2}, DataType::kFLOAT32, GetDevice()); + + auto swiglu_fn = std::make_shared(); + swiglu_fn->Apply({input}); + auto grad_inputs = swiglu_fn->Backward({grad_output}); + ASSERT_EQ(grad_inputs.size(), 1); + EXPECT_EQ(grad_inputs[0]->Dtype(), DataType::kFLOAT32); + + std::vector expected_grad(4); + for (int64_t col = 0; col < 2; ++col) { + const float gate = input_values[col]; + const float up = input_values[2 + col]; + const float grad = grad_values[col]; + const float sigmoid = 1.0f / (1.0f + std::exp(-gate)); + expected_grad[col] = grad * up * sigmoid * (1.0f + gate * (1.0f - sigmoid)); + expected_grad[2 + col] = grad * gate * sigmoid; + } + test::ExpectTensorNear(grad_inputs[0], expected_grad, 2e-3f); +} + TEST_P(AutogradElementwiseBackwardTest, AddBackward) { auto a = std::make_shared(std::vector{2, 3}, DataType::kFLOAT32, GetDevice(), true); a->Fill(1.0f); diff --git a/tests/transformer/test_transformer_architecture.cc b/tests/transformer/test_transformer_architecture.cc index d4a6efc2..4cec471d 100644 --- a/tests/transformer/test_transformer_architecture.cc +++ b/tests/transformer/test_transformer_architecture.cc @@ -88,7 +88,11 @@ TEST_P(TransformerModuleTest, SwiGLUMLP) { auto mlp = std::make_shared(config); mlp->To(GetDevice()); - EXPECT_EQ(mlp->Parameters().size(), 3); + EXPECT_EQ(mlp->Parameters().size(), 2); + auto state = mlp->StateDict(); + ASSERT_TRUE(state.contains("c_fc.weight")); + EXPECT_FALSE(state.contains("c_fc2.weight")); + EXPECT_EQ(state.at("c_fc.weight")->Dims(), (std::vector{512, config.n_embd})); auto input = std::make_shared(std::vector{2, 8, 64}, DataType::kFLOAT32, GetDevice()); auto output = (*mlp)({input}); @@ -241,10 +245,9 @@ TEST_P(TransformerModuleTest, MoELayerTop2SwiGLU) { auto state = moe->StateDict(); ASSERT_TRUE(state.contains("experts.expert_0.c_fc.weight")); - ASSERT_TRUE(state.contains("experts.expert_0.c_fc2.weight")); + EXPECT_FALSE(state.contains("experts.expert_0.c_fc2.weight")); ASSERT_TRUE(state.contains("experts.expert_0.c_proj.weight")); - EXPECT_EQ(state.at("experts.expert_0.c_fc.weight")->Dims(), (std::vector{48, config.n_embd})); - EXPECT_EQ(state.at("experts.expert_0.c_fc2.weight")->Dims(), (std::vector{48, config.n_embd})); + EXPECT_EQ(state.at("experts.expert_0.c_fc.weight")->Dims(), (std::vector{96, config.n_embd})); EXPECT_EQ(state.at("experts.expert_0.c_proj.weight")->Dims(), (std::vector{config.n_embd, 48})); }