MCPcopy Create free account
hub / github.com/WasmEdge/WasmEdge / run

Method run

plugins/wasi_nn/MLX/model/whisper/decoding.cpp:710–862  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

708}
709
710std::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 }

Callers 1

decodeFunction · 0.45

Calls 12

detectLanguageFunction · 0.85
takeFunction · 0.85
compressionRatioFunction · 0.85
resizeMethod · 0.80
getEotMethod · 0.80
rankMethod · 0.80
decodeMethod · 0.80
eraseMethod · 0.80
resetMethod · 0.45
dataMethod · 0.45
sizeMethod · 0.45
finalizeMethod · 0.45

Tested by

no test coverage detected