| 645 | } |
| 646 | |
| 647 | std::tuple<mx::array, mx::array, mx::array> |
| 648 | DecodingTask::mainLoop(const mx::array &AudioFeatures, |
| 649 | const mx::array &Tokens) { |
| 650 | |
| 651 | int NBatch = Tokens.shape(0); |
| 652 | mx::array CurrentTokens = Tokens; |
| 653 | mx::array SumLogprobs = mx::zeros({NBatch}, mx::float32); |
| 654 | bool Completed = false; |
| 655 | |
| 656 | auto StepFunction = [&](const mx::array &Inputs, const mx::array &AudioFeats, |
| 657 | const mx::array &TokSeq, const mx::array &SumLogp) |
| 658 | -> std::tuple<mx::array, bool, mx::array, mx::array> { |
| 659 | mx::array PreLogits = Inference->logits(Inputs, AudioFeats); |
| 660 | mx::array Logits = take(PreLogits, PreLogits.shape(1) - 1, 1); |
| 661 | for (const auto &Filter : LogitFilters) { |
| 662 | Logits = Filter->apply(Logits, TokSeq); |
| 663 | } |
| 664 | auto [NextTokens, CompletedFlag, NextSumLogprobs] = |
| 665 | Decoder->update(TokSeq, Logits, SumLogp); |
| 666 | return std::make_tuple(NextTokens, CompletedFlag, NextSumLogprobs, |
| 667 | PreLogits); |
| 668 | }; |
| 669 | |
| 670 | auto [NextTokens, CompletedFlag, NextSumLogprobs, PreLogits] = |
| 671 | StepFunction(CurrentTokens, AudioFeatures, CurrentTokens, SumLogprobs); |
| 672 | |
| 673 | CurrentTokens = NextTokens; |
| 674 | SumLogprobs = NextSumLogprobs; |
| 675 | Completed = CompletedFlag; |
| 676 | |
| 677 | mx::array NoSpeechProbs = mx::zeros({NBatch}, mx::float32); |
| 678 | if (Tokenizer->getNoSpeech() != -1) { |
| 679 | auto ProbsAtSot = mx::softmax(mx::take(PreLogits, SotIndex, 1), -1); |
| 680 | NoSpeechProbs = mx::take(ProbsAtSot, Tokenizer->getNoSpeech(), 1); |
| 681 | } else { |
| 682 | NoSpeechProbs = mx::full({NBatch}, std::numeric_limits<float>::quiet_NaN(), |
| 683 | mx::float32); |
| 684 | } |
| 685 | |
| 686 | mx::eval(CurrentTokens, SumLogprobs, NoSpeechProbs); |
| 687 | for (int I = 1; I < SampleLen; ++I) { |
| 688 | mx::array Inputs = |
| 689 | take(CurrentTokens, mx::array({CurrentTokens.shape(1) - 1}), 1); |
| 690 | |
| 691 | if (CurrentTokens.shape(-1) > NCtx) { |
| 692 | break; |
| 693 | } |
| 694 | auto [NextToks, NextCompleted, NextSumLogp, _] = |
| 695 | StepFunction(Inputs, AudioFeatures, CurrentTokens, SumLogprobs); |
| 696 | mx::eval(NextToks, NextSumLogp); |
| 697 | |
| 698 | if (Completed) { |
| 699 | break; |
| 700 | } |
| 701 | |
| 702 | CurrentTokens = NextToks; |
| 703 | Completed = NextCompleted; |
| 704 | SumLogprobs = NextSumLogp; |
nothing calls this directly
no test coverage detected