| 196 | } |
| 197 | |
| 198 | HeartMuLaFrameEmbeddingInputs audio_frame_embedding_inputs( |
| 199 | const std::vector<int32_t> & frame_tokens, |
| 200 | int64_t batch, |
| 201 | int64_t codebook, |
| 202 | const HeartMuLaConfig & config) { |
| 203 | if (static_cast<int64_t>(frame_tokens.size()) != batch) { |
| 204 | throw std::runtime_error("HeartMuLa audio embedding token batch mismatch"); |
| 205 | } |
| 206 | HeartMuLaFrameEmbeddingInputs inputs; |
| 207 | inputs.batch_size = batch; |
| 208 | inputs.steps = 1; |
| 209 | const size_t codebooks = static_cast<size_t>(config.audio_num_codebooks); |
| 210 | inputs.audio_token_ids.assign(static_cast<size_t>(batch) * codebooks, 0); |
| 211 | inputs.audio_mask.assign(static_cast<size_t>(batch) * codebooks, 0.0F); |
| 212 | inputs.text_token_ids.assign(static_cast<size_t>(batch), 0); |
| 213 | inputs.text_cond_mask.assign(static_cast<size_t>(batch), 0.0F); |
| 214 | inputs.text_uncond_mask.assign(static_cast<size_t>(batch), 0.0F); |
| 215 | for (int64_t b = 0; b < batch; ++b) { |
| 216 | for (int64_t q = 0; q < config.audio_num_codebooks; ++q) { |
| 217 | const size_t index = static_cast<size_t>(b * config.audio_num_codebooks + q); |
| 218 | const int32_t raw = q == codebook ? frame_tokens[static_cast<size_t>(b)] : 0; |
| 219 | inputs.audio_token_ids[index] = raw + static_cast<int32_t>(q * config.audio_vocab_size); |
| 220 | inputs.audio_mask[index] = q == codebook ? 1.0F : 0.0F; |
| 221 | } |
| 222 | } |
| 223 | return inputs; |
| 224 | } |
| 225 | |
| 226 | HeartMuLaFrameEmbeddingInputs next_frame_embedding_inputs( |
| 227 | const std::vector<int32_t> & frame_tokens, |
no test coverage detected