| 148 | } |
| 149 | |
| 150 | HeartMuLaFrameEmbeddingInputs prompt_embedding_inputs( |
| 151 | const HeartMuLaPromptEncoding & encoding, |
| 152 | const HeartMuLaConfig & config, |
| 153 | bool use_cfg) { |
| 154 | HeartMuLaFrameEmbeddingInputs inputs; |
| 155 | inputs.batch_size = encoding.batch_size; |
| 156 | inputs.steps = encoding.prompt_len; |
| 157 | const size_t batch = static_cast<size_t>(encoding.batch_size); |
| 158 | const size_t steps = static_cast<size_t>(encoding.prompt_len); |
| 159 | const size_t lanes = static_cast<size_t>(encoding.parallel_number); |
| 160 | const size_t codebooks = static_cast<size_t>(config.audio_num_codebooks); |
| 161 | inputs.audio_token_ids.assign(batch * steps * codebooks, 0); |
| 162 | inputs.text_token_ids.assign(batch * steps, 0); |
| 163 | inputs.audio_mask.assign(batch * steps * codebooks, 0.0F); |
| 164 | inputs.text_cond_mask.assign(batch * steps, 0.0F); |
| 165 | inputs.text_uncond_mask.assign(batch * steps, 0.0F); |
| 166 | const size_t actual_batch = use_cfg ? batch / 2 : batch; |
| 167 | for (size_t b = 0; b < batch; ++b) { |
| 168 | const bool uncond = use_cfg && b >= actual_batch; |
| 169 | for (size_t t = 0; t < steps; ++t) { |
| 170 | for (size_t q = 0; q < codebooks; ++q) { |
| 171 | const size_t src = prompt_flat(b, t, q, steps, lanes); |
| 172 | const size_t dst = (b * steps + t) * codebooks + q; |
| 173 | inputs.audio_token_ids[dst] = static_cast<int32_t>( |
| 174 | encoding.tokens[src] + static_cast<int64_t>(q) * config.audio_vocab_size); |
| 175 | inputs.audio_mask[dst] = encoding.tokens_mask[src] != 0U ? 1.0F : 0.0F; |
| 176 | } |
| 177 | const size_t text_src = prompt_flat(b, t, lanes - 1, steps, lanes); |
| 178 | inputs.text_token_ids[b * steps + t] = static_cast<int32_t>(encoding.tokens[text_src]); |
| 179 | const float text_mask = encoding.tokens_mask[text_src] != 0U ? 1.0F : 0.0F; |
| 180 | inputs.text_cond_mask[b * steps + t] = uncond ? 0.0F : text_mask; |
| 181 | inputs.text_uncond_mask[b * steps + t] = uncond ? text_mask : 0.0F; |
| 182 | } |
| 183 | } |
| 184 | inputs.apply_muq = true; |
| 185 | inputs.muq_row = encoding.muq_idx.empty() ? 0 : encoding.muq_idx.front(); |
| 186 | inputs.muq_embed = encoding.muq_embed; |
| 187 | inputs.muq_cond_mask.assign(batch, 1.0F); |
| 188 | inputs.muq_uncond_mask.assign(batch, 0.0F); |
| 189 | if (use_cfg) { |
| 190 | for (size_t b = actual_batch; b < batch; ++b) { |
| 191 | inputs.muq_cond_mask[b] = 0.0F; |
| 192 | inputs.muq_uncond_mask[b] = 1.0F; |
| 193 | } |
| 194 | } |
| 195 | return inputs; |
| 196 | } |
| 197 | |
| 198 | HeartMuLaFrameEmbeddingInputs audio_frame_embedding_inputs( |
| 199 | const std::vector<int32_t> & frame_tokens, |
no test coverage detected