MCPcopy Create free account
hub / github.com/0xShug0/audio.cpp / encode_prompt

Method encode_prompt

src/models/heartmula/tokenizer_text.cpp:81–129  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

79}
80
81HeartMuLaPromptEncoding HeartMuLaTextTokenizer::encode_prompt(const HeartMuLaPromptRequest & request) const {
82 const auto & config = impl_->assets->generation_config;
83 const auto & model_config = impl_->assets->mula_config;
84 auto tags_ids = encode(normalize_tags(request.tags));
85 auto lyrics_ids = encode(unicode_lowercase(request.lyrics));
86 add_bos_eos(
87 tags_ids,
88 static_cast<int32_t>(config.text_bos_id),
89 static_cast<int32_t>(config.text_eos_id),
90 "tags");
91 add_bos_eos(
92 lyrics_ids,
93 static_cast<int32_t>(config.text_bos_id),
94 static_cast<int32_t>(config.text_eos_id),
95 "lyrics");
96
97 HeartMuLaPromptEncoding encoding;
98 encoding.batch_size = request.options.guidance_scale != 1.0F ? 2 : 1;
99 encoding.prompt_len = static_cast<int64_t>(tags_ids.size() + 1 + lyrics_ids.size());
100 encoding.parallel_number = model_config.audio_num_codebooks + 1;
101 encoding.tags_ids = tags_ids;
102 encoding.lyrics_ids = lyrics_ids;
103
104 const auto batch_size = static_cast<size_t>(encoding.batch_size);
105 const auto prompt_len = static_cast<size_t>(encoding.prompt_len);
106 const auto parallel_number = static_cast<size_t>(encoding.parallel_number);
107 const auto total_token_values = batch_size * prompt_len * parallel_number;
108 encoding.tokens.assign(total_token_values, config.empty_id);
109 encoding.tokens_mask.assign(total_token_values, 0U);
110 encoding.muq_embed.assign(batch_size * static_cast<size_t>(model_config.muq_dim), 0.0F);
111 encoding.muq_idx.assign(batch_size, static_cast<int64_t>(tags_ids.size()));
112 encoding.pos.resize(batch_size * prompt_len);
113
114 const size_t text_lane = parallel_number - 1;
115 for (size_t b = 0; b < batch_size; ++b) {
116 for (size_t i = 0; i < tags_ids.size(); ++i) {
117 encoding.tokens[(b * prompt_len + i) * parallel_number + text_lane] = tags_ids[i];
118 }
119 const size_t lyrics_row_offset = tags_ids.size() + 1;
120 for (size_t i = 0; i < lyrics_ids.size(); ++i) {
121 encoding.tokens[(b * prompt_len + lyrics_row_offset + i) * parallel_number + text_lane] = lyrics_ids[i];
122 }
123 for (size_t row = 0; row < prompt_len; ++row) {
124 encoding.tokens_mask[(b * prompt_len + row) * parallel_number + text_lane] = 1U;
125 encoding.pos[b * prompt_len + row] = static_cast<int64_t>(row);
126 }
127 }
128 return encoding;
129}
130
131} // namespace engine::models::heartmula

Callers 1

Calls 6

normalize_tagsFunction · 0.85
unicode_lowercaseFunction · 0.85
add_bos_eosFunction · 0.85
assignMethod · 0.80
sizeMethod · 0.45
resizeMethod · 0.45

Tested by

no test coverage detected