| 79 | } |
| 80 | |
| 81 | HeartMuLaPromptEncoding 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 |
no test coverage detected