| 708 | } |
| 709 | |
| 710 | std::vector<DecodingResult> DecodingTask::run(const mx::array &Mel) { |
| 711 | Inference->reset(); |
| 712 | Decoder->reset(); |
| 713 | int NAudio = Mel.shape(0); |
| 714 | |
| 715 | mx::array AudioFeatures = getAudioFeatures(Mel); |
| 716 | |
| 717 | mx::array Tokens = |
| 718 | mx::array(InitialTokens.data(), {static_cast<int>(InitialTokens.size())}, |
| 719 | mx::int32); |
| 720 | Tokens = mx::broadcast_to(Tokens, |
| 721 | {NAudio, static_cast<int>(InitialTokens.size())}); |
| 722 | auto [Languages, LangProbs] = detectLanguage(AudioFeatures, Tokens); |
| 723 | |
| 724 | if (Options.Task == "lang_id") { |
| 725 | std::vector<DecodingResult> Results; |
| 726 | for (int I = 0; I < NAudio; ++I) { |
| 727 | DecodingResult Result; |
| 728 | Result.AudioFeatures = mx::take(AudioFeatures, mx::array({I}), 0); |
| 729 | Result.Language = Languages[I]; |
| 730 | if (LangProbs) { |
| 731 | Result.LanguageProbs = (*LangProbs)[I]; |
| 732 | } |
| 733 | Results.push_back(Result); |
| 734 | } |
| 735 | return Results; |
| 736 | } |
| 737 | |
| 738 | if (NGroup > 1) { |
| 739 | // tokens = tokens[:, None, :] |
| 740 | Tokens = mx::expand_dims(Tokens, 1); |
| 741 | |
| 742 | // tokens = mx.broadcast_to(tokens, [n_audio, self.n_group, |
| 743 | // len(self.initial_tokens)]) |
| 744 | std::vector<int> NewShape = {NAudio, NGroup, |
| 745 | static_cast<int>(InitialTokens.size())}; |
| 746 | Tokens = mx::broadcast_to(Tokens, NewShape); |
| 747 | |
| 748 | // tokens = tokens.reshape(n_audio * self.n_group, len(self.initial_tokens)) |
| 749 | Tokens = mx::reshape( |
| 750 | Tokens, {NAudio * NGroup, static_cast<int>(InitialTokens.size())}); |
| 751 | } |
| 752 | // Call the main sampling loop |
| 753 | auto [TokensResult, SumLogprobs, NoSpeechProbs] = |
| 754 | mainLoop(AudioFeatures, Tokens); |
| 755 | |
| 756 | // Reshape the tensors to have (n_audio, n_group) as the first two dimensions |
| 757 | AudioFeatures = |
| 758 | mx::take(AudioFeatures, mx::arange(0, AudioFeatures.shape(0), NGroup), 0); |
| 759 | NoSpeechProbs = |
| 760 | mx::take(NoSpeechProbs, mx::arange(0, NoSpeechProbs.shape(0), NGroup), 0); |
| 761 | |
| 762 | // Ensure shapes are consistent |
| 763 | if (AudioFeatures.shape(0) != NoSpeechProbs.shape(0) || |
| 764 | AudioFeatures.shape(0) != NAudio) { |
| 765 | throw std::runtime_error( |
| 766 | "Inconsistent audio features and no_speech_probs shapes"); |
| 767 | } |
no test coverage detected