Skip to content
Merged
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
2 changes: 2 additions & 0 deletions sherpa-onnx/csrc/offline-canary-model-meta-data.h
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@ struct OfflineCanaryModelMetaData {
int32_t vocab_size;
int32_t subsampling_factor = 8;
int32_t feat_dim = 120;
int32_t num_decoder_layers = 6;
int32_t decoder_hidden_size = 1024;
std::string normalize_type;
std::unordered_map<std::string, int32_t> lang2id;
};
Expand Down
13 changes: 10 additions & 3 deletions sherpa-onnx/csrc/offline-canary-model.cc
Original file line number Diff line number Diff line change
Expand Up @@ -117,11 +117,13 @@ class OfflineCanaryModel::Impl {
}

std::vector<Ort::Value> GetInitialDecoderStates() {
std::array<int64_t, 3> shape{1, 0, 1024};
int32_t num_layers = meta_.num_decoder_layers;
int64_t hidden_size = meta_.decoder_hidden_size;
std::array<int64_t, 3> shape{1, 0, hidden_size};

std::vector<Ort::Value> ans;
ans.reserve(6);
for (int32_t i = 0; i < 6; ++i) {
ans.reserve(num_layers);
for (int32_t i = 0; i < num_layers; ++i) {
Ort::Value state = Ort::Value::CreateTensor<float>(
Allocator(), shape.data(), shape.size());

Expand Down Expand Up @@ -178,6 +180,11 @@ class OfflineCanaryModel::Impl {
"normalize_type");
SHERPA_ONNX_READ_META_DATA(meta_.subsampling_factor, "subsampling_factor");
SHERPA_ONNX_READ_META_DATA(meta_.feat_dim, "feat_dim");

SHERPA_ONNX_READ_META_DATA_WITH_DEFAULT(meta_.num_decoder_layers,
"num_decoder_layers", 6);
SHERPA_ONNX_READ_META_DATA_WITH_DEFAULT(meta_.decoder_hidden_size,
"decoder_hidden_size", 1024);
}

void InitDecoder(void *model_data, size_t model_data_length) {
Expand Down