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
66 changes: 20 additions & 46 deletions example/gpt2/checkpoint_loader.cc
Original file line number Diff line number Diff line change
Expand Up @@ -163,44 +163,38 @@ std::shared_ptr<nn::TransformerModel> LoadFromLLMC(const std::string &filepath)
}

// transformer.h.{i}.ln_1.weight
int local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor
= state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName,
nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index),
nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamWeightName)];
auto &tensor = state_dict[std::format(
"{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName,
std::to_string(idx), nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamWeightName)];
ReadVectorAllFloat(ifs, static_cast<float *>(tensor->DataPtr()), n_embd);
++local_layer_index;
} else {
size_t ln_1_w_bytes = n_embd * sizeof(float);
ifs.seekg(ln_1_w_bytes, std::ios::cur);
}
}

// transformer.h.{i}.ln_1.bias
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName,
nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index),
nn::TransformerChunk::kHLayerName, std::to_string(idx),
nn::TransformerLayer::kLn1LayerName, nn::LayerNorm::kParamBiasName)];
ReadVectorAllFloat(ifs, static_cast<float *>(tensor->DataPtr()), n_embd);
++local_layer_index;
} else {
size_t ln_1_b_bytes = n_embd * sizeof(float);
ifs.seekg(ln_1_b_bytes, std::ios::cur);
}
}

// transformer.h.{i}.attn.c_attn.weight (ColumnParallelLinear, but actually applies on "rows")
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format(
"{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName,
std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName,
nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamWeightName)];
std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName,
nn::parallel::ColumnParallelLinear::kParamWeightName)];
// NOTE(zbl): In the .bin model file, Q/K/V is concated along last dim,
// i.e. [Q|K|V].T = [q1|q2|...|qn|k1|k2|...|kn|v1|v2|...|vn].T
// However, each tp_rank needs to get [q_i|k_i|v_i].T, so we need to jump and read them
Expand Down Expand Up @@ -229,21 +223,19 @@ std::shared_ptr<nn::TransformerModel> LoadFromLLMC(const std::string &filepath)
/*rows=*/rows_all, /*cols=*/cols_all,
/*row_start=*/2 * n_embd + tp_rank * local_C, /*row_cnt=*/local_C);

++local_layer_index;
} else {
size_t c_attn_w_bytes = qkv_out * n_embd * sizeof(float);
ifs.seekg(c_attn_w_bytes, std::ios::cur);
}
}

// transformer.h.{i}.attn.c_attn.bias (ColumnParallelLinear)
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format(
"{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName,
std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName,
nn::CausalSelfAttention::kCAttnLayerName, nn::parallel::ColumnParallelLinear::kParamBiasName)];
std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCAttnLayerName,
nn::parallel::ColumnParallelLinear::kParamBiasName)];
// NOTE(zbl): Same as c_attn.weight, the bias for Q/K/V is concated
// i.e. [Q|K|V] = [q1|q2|...|qn|k1|k2|...|kn|v1|v2|...|vn]
// However, each tp_rank needs to get [q_i|k_i|v_i], so we need to jump and read them
Expand Down Expand Up @@ -271,136 +263,118 @@ std::shared_ptr<nn::TransformerModel> LoadFromLLMC(const std::string &filepath)
/*len=*/len_all,
/*start=*/2 * n_embd + tp_rank * local_C, /*cnt=*/local_C);

++local_layer_index;
} else {
size_t c_attn_b_bytes = qkv_out * sizeof(float);
ifs.seekg(c_attn_b_bytes, std::ios::cur);
}
}

// transformer.h.{i}.attn.c_proj.weight (RowParallelLinear, but actually applies on "columns")
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format(
"{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName,
std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName,
nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamWeightName)];
std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName,
nn::parallel::RowParallelLinear::kParamWeightName)];
ReadMatrixColShardFloat(ifs, static_cast<float *>(tensor->DataPtr()), n_embd, n_embd, tp_rank * in_pp,
in_pp);
++local_layer_index;
} else {
size_t c_proj_w_bytes = n_embd * n_embd * sizeof(float);
ifs.seekg(c_proj_w_bytes, std::ios::cur);
}
}

// transformer.h.{i}.attn.c_proj.bias (RowParallelLinear, no shard on bias)
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format(
"{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName,
std::to_string(local_layer_index), nn::TransformerLayer::kAttnLayerName,
nn::CausalSelfAttention::kCProjLayerName, nn::parallel::RowParallelLinear::kParamBiasName)];
std::to_string(idx), nn::TransformerLayer::kAttnLayerName, nn::CausalSelfAttention::kCProjLayerName,
nn::parallel::RowParallelLinear::kParamBiasName)];
ReadVectorAllFloat(ifs, static_cast<float *>(tensor->DataPtr()), n_embd);
++local_layer_index;
} else {
size_t c_proj_b_bytes = n_embd * sizeof(float);
ifs.seekg(c_proj_b_bytes, std::ios::cur);
}
}

// transformer.h.{i}.ln_2.weight
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor
= state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName,
nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index),
nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamWeightName)];
auto &tensor = state_dict[std::format(
"{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName, nn::TransformerChunk::kHLayerName,
std::to_string(idx), nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamWeightName)];
ReadVectorAllFloat(ifs, static_cast<float *>(tensor->DataPtr()), n_embd);
++local_layer_index;
} else {
size_t ln_2_w_bytes = n_embd * sizeof(float);
ifs.seekg(ln_2_w_bytes, std::ios::cur);
}
}

// transformer.h.{i}.ln_2.bias
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName,
nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index),
nn::TransformerChunk::kHLayerName, std::to_string(idx),
nn::TransformerLayer::kLn2LayerName, nn::LayerNorm::kParamBiasName)];
ReadVectorAllFloat(ifs, static_cast<float *>(tensor->DataPtr()), n_embd);
++local_layer_index;
} else {
size_t ln_2_b_bytes = n_embd * sizeof(float);
ifs.seekg(ln_2_b_bytes, std::ios::cur);
}
}

// transformer.h.{i}.mlp.c_fc.weight (ColumnParallelLinear, but actually applies on "rows")
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName,
nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index),
nn::TransformerChunk::kHLayerName, std::to_string(idx),
nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName,
nn::parallel::ColumnParallelLinear::kParamWeightName)];
ReadMatrixRowShardFloat(ifs, static_cast<float *>(tensor->DataPtr()), fc_out, n_embd, fc_start, fc_pp);
++local_layer_index;
} else {
size_t c_fc_w_bytes = fc_out * n_embd * sizeof(float);
ifs.seekg(c_fc_w_bytes, std::ios::cur);
}
}

// transformer.h.{i}.mlp.c_fc.bias (ColumnParallelLinear)
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName,
nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index),
nn::TransformerChunk::kHLayerName, std::to_string(idx),
nn::TransformerLayer::kMlpLayerName, nn::MLP::kCFcLayerName,
nn::parallel::ColumnParallelLinear::kParamBiasName)];
ReadVectorShardFloat(ifs, static_cast<float *>(tensor->DataPtr()), fc_out, fc_start, fc_pp);
++local_layer_index;
} else {
size_t c_fc_b_bytes = fc_out * sizeof(float);
ifs.seekg(c_fc_b_bytes, std::ios::cur);
}
}

// transformer.h.{i}.mlp.c_proj.weight (RowParallelLinear, but actually applies on "columns")
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName,
nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index),
nn::TransformerChunk::kHLayerName, std::to_string(idx),
nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName,
nn::parallel::RowParallelLinear::kParamWeightName)];
ReadMatrixColShardFloat(ifs, static_cast<float *>(tensor->DataPtr()), n_embd, fc_out, tp_rank * in4_pp,
in4_pp);
++local_layer_index;
} else {
size_t c_proj_w_bytes = fc_out * n_embd * sizeof(float);
ifs.seekg(c_proj_w_bytes, std::ios::cur);
}
}

// transformer.h.{i}.mlp.c_proj.bias (RowParallelLinear, no shard on bias)
local_layer_index = 0;
for (int idx = 0; idx < n_layer; ++idx) {
if (owned_layers[idx]) {
auto &tensor = state_dict[std::format("{}.{}.{}.{}.{}.{}", nn::TransformerModel::kTransformerModelName,
nn::TransformerChunk::kHLayerName, std::to_string(local_layer_index),
nn::TransformerChunk::kHLayerName, std::to_string(idx),
nn::TransformerLayer::kMlpLayerName, nn::MLP::kCProjLayerName,
nn::parallel::RowParallelLinear::kParamBiasName)];
ReadVectorAllFloat(ifs, static_cast<float *>(tensor->DataPtr()), n_embd);
++local_layer_index;
} else {
size_t c_proj_b_bytes = n_embd * sizeof(float);
ifs.seekg(c_proj_b_bytes, std::ios::cur);
Expand Down
6 changes: 4 additions & 2 deletions example/gpt2/main.cc
Original file line number Diff line number Diff line change
Expand Up @@ -357,7 +357,6 @@ void Train(const nn::parallel::Rank &rank) {
} else {
optimizer = optimizer_creator(named_parameters);
}

const int64_t lr_decay_iters = FLAGS_lr_decay_iters > 0 ? FLAGS_lr_decay_iters : FLAGS_num_iteration;
TrainingLRSchedulerConfig sched_config;
sched_config.lr = static_cast<float>(FLAGS_learning_rate);
Expand Down Expand Up @@ -422,11 +421,13 @@ void Train(const nn::parallel::Rank &rank) {
.n_head = model_config.n_head,
.n_kv_head = model_config.n_kv_head,
.n_embd = model_config.n_embd,
.vocab_size = model_config.vocab_size,
.original_vocab_size = model_config.original_vocab_size,
.padded_vocab_size = model_config.vocab_size,
.ddp_size = ddp_world_size,
.tp_size = tp_world_size,
.sp_size = sp_world_size,
.pp_size = pp_world_size,
.vpp_size = static_cast<int>(FLAGS_virtual_pipeline_parallel),
.checkpoint_root_dir = FLAGS_save,
.max_checkpoint_keep = FLAGS_max_checkpoint_keep,
.rank = rank,
Expand Down Expand Up @@ -516,6 +517,7 @@ void Train(const nn::parallel::Rank &rank) {
LOG(INFO) << "Rank " << rank.GlobalRank() << ": finish backward";
}

nn::parallel::FinalizeModelGrads({model});
optimizer->Step();
if (scheduler) {
scheduler->Step();
Expand Down
Loading
Loading