#include "streaming.hpp" #include "tokenizer.hpp" #include "mel.hpp" #include #include namespace pk { StreamingSession::StreamingSession(const ModelLoader& ml, const std::string& target_lang) : ml_(ml), enc_(ml), pred_(ml), joint_(ml), prompt_(ml) { const ParakeetConfig& cfg = ml.config(); d_model_ = (int)cfg.d_model; blank_id_ = (int)cfg.blank_id; // Resolve the language prompt index for multilingual (nemotron) models. The // one-hot is constant over time (one language per utterance), so the prompt // projection is applied per chunk in feed_mel_chunk using this fixed index. // Non-prompt models leave prompt_index_ = -1 and skip prompt_.apply(). if (cfg.prompt.present) { // Empty target_lang -> the model default; an unknown locale THROWS // std::runtime_error (same message as Model::resolve_prompt_index), so a // typo (e.g. --lang xx) fails loudly instead of silently mis-transcribing. // Matches the offline path and the parakeet_capi_stream_begin_lang // contract (NULL + ctx last_error on an unknown locale). prompt_index_ = cfg.prompt.resolve_index_or_throw(target_lang); } // Greedy max symbols per frame, from model metadata (NeMo default 10); // matches the offline pk::transcribe path in model.cpp. max_symbols_ = (int)cfg.max_symbols; assert(joint_.num_durations() == 0 && "StreamingSession is RNN-T only (no TDT durations)"); // Resolve the / special token ids from the tokenizer pieces (do // NOT hardcode 1024/1025; read them from the loaded vocab). const auto& pieces = cfg.tokenizer_pieces; for (int i = 0; i < (int)pieces.size(); ++i) { if (pieces[i] == "") eou_id_ = i; else if (pieces[i] == "") eob_id_ = i; } // Per-encoder-frame stride in seconds: each encoder output frame spans // hop_length * subsampling_factor input samples (the subsampling downsamples // the mel-frame time axis by subsampling_factor). time = frame * stride. const double hop = (double)cfg.hop_length; const double sub = (double)(cfg.subsampling_factor ? cfg.subsampling_factor : 1); const double sr = (double)(cfg.sample_rate ? cfg.sample_rate : 16000); frame_sec_ = (hop * sub) / sr; frame_sec_f_ = (float)frame_sec_; reset(); } void StreamingSession::reset() { enc_.reset(); state_ = rnnt_decode_init(pred_); // eou_id_/eob_id_/frame_sec_ are resolved once in the constructor. enc_frame_ = 0; non_special_.clear(); text_.clear(); text_taken_ = 0; last_chunk_had_eou_ = false; events_.clear(); word_tokens_.clear(); words_.clear(); words_finalized_ = 0; words_taken_ = 0; eou_closed_words_ = 0; } void StreamingSession::process_emitted(const std::vector& emitted) { last_chunk_had_eou_ = false; bool text_changed = false; for (int32_t tok : emitted) { if (tok == eou_id_ || tok == eob_id_) { // EOU/EOB: surface as an event, do NOT add to the text. EouEvent ev; ev.token = tok; ev.is_eob = (tok == eob_id_); // The frame index is recorded into events_ by feed/finalize (which // knows the per-token frame); we patch it there. Here we only know // the token, so push a placeholder the caller will overwrite. ev.encoder_frame = enc_frame_; // refined by caller via per-token frame ev.time_sec = ev.encoder_frame * frame_sec_; events_.push_back(ev); last_chunk_had_eou_ = true; } else { non_special_.push_back(tok); text_changed = true; } } if (text_changed) { text_ = detokenize(ml_.config().tokenizer_pieces, non_special_); } } std::vector StreamingSession::feed_mel_chunk(const std::vector& mel_chunk, int n_frames, bool is_last) { // 1. Encoder step: the chunk's valid encoder frames, row-major [valid, d_model] // (d_model fastest) — exactly the orientation rnnt_decode_frames expects. int n_valid = 0; std::vector enc_frames = enc_.step(mel_chunk, n_frames, is_last, n_valid); if (n_valid <= 0) { last_chunk_had_eou_ = false; return {}; } // 1b. Prompt conditioning (nemotron multilingual): project the chunk's // encoder frames through prompt_kernel for the resolved language before // the RNN-T decode. The one-hot is constant over time, so applying it // per chunk is exact (== the offline forward's single application). // prompt_.apply() wants channels-first [d_model, valid]; enc_.step gives // time-major [valid, d_model], so transpose in, apply, transpose back. // No-op for non-prompt models (prompt_.present()==false): enc_frames is // left byte-identical. if (prompt_.present()) { std::vector chunk_cf((size_t)d_model_ * n_valid); // [d_model, valid] for (int t = 0; t < n_valid; ++t) for (int c = 0; c < d_model_; ++c) chunk_cf[(size_t)c * n_valid + t] = enc_frames[(size_t)t * d_model_ + c]; std::vector projected; prompt_.apply(chunk_cf, d_model_, n_valid, prompt_index_, projected); // [d_model, valid] for (int t = 0; t < n_valid; ++t) for (int c = 0; c < d_model_; ++c) enc_frames[(size_t)t * d_model_ + c] = projected[(size_t)c * n_valid + t]; } // 2. RNN-T greedy over the new encoder frames, carrying the decoder state // across chunks (do NOT reset). Appends to state_.hyp and returns the ids // emitted in this chunk, with their LOCAL frame index in [0, n_valid) and // per-token TokenInfo (LOCAL frame, max_prob conf, span==1). const int base_frame = enc_frame_; std::vector local_frames; std::vector chunk_tokens; std::vector emitted = rnnt_decode_frames(pred_, joint_, enc_frames, n_valid, d_model_, state_, blank_id_, max_symbols_, &local_frames, &chunk_tokens); enc_frame_ += n_valid; // 3. Update text + EOU events; refine each new event's absolute frame index // from the per-token local frame the decoder reported. const size_t prev_events = events_.size(); process_emitted(emitted); // Re-walk emitted to assign the correct absolute frame to each new event, // and accumulate NON-special tokens (absolute frame) for word grouping. size_t evi = prev_events; size_t eou_word_tokens = 0; for (size_t i = 0; i < emitted.size(); ++i) { if (emitted[i] == eou_id_ || emitted[i] == eob_id_) { const int abs_frame = base_frame + (int)local_frames[i]; events_[evi].encoder_frame = abs_frame; events_[evi].time_sec = abs_frame * frame_sec_; ++evi; // Tokens are walked in emission order, so this is the word-token // count the closes. eou_word_tokens = word_tokens_.size(); } else { TokenInfo ti = chunk_tokens[i]; ti.frame += base_frame; // local -> absolute encoder frame word_tokens_.push_back(ti); } } // Regroup the accumulated tokens; words before the last (still-open) one are // final and become available to drain_words(), as are the words closed by an // in this chunk. regroup_words(/*flush_all=*/false, eou_word_tokens); // 4. End-of-utterance reset. The realtime EOU model is trained to emit // (end of utterance) / (backchannel) and have the decoder START THE // NEXT UTTERANCE FROM A FRESH STATE — exactly what NeMo's reference // streaming driver does (examples/voice_agent .../nemo/streaming_asr.py // NemoStreamingASRService.transcribe -> reset_state() when / // appears in the chunk text). Without this the prediction net stays // conditioned on the just-emitted and the joint scores blank on every // subsequent frame, so the stream goes silent after the first utterance // (issue #13). We reset the carried RNN-T decoder state — LSTM h/c to zero // and last_token back to SOS — for the next chunk. We deliberately reset // ONLY the decoder, not the StreamingEncoder cache: NeMo's reset_state also // drops the encoder cache, but that was verified to be a no-op for the // decoded tokens (decoder-only reset == NeMo's full reset_state byte-for- // byte on multi-utterance clips), so the validated streaming-encoder path is // left untouched. enc_frame_ keeps running so timestamps stay absolute // in the clip, and state_.hyp keeps the full token record across utterances. if (last_chunk_had_eou_) { state_.state = pred_.zero_state(); state_.last_token = -1; // SOS sentinel (nothing emitted yet) state_.have_token = false; } return emitted; } std::string StreamingSession::finalize() { // The end-of-stream tail is flushed by the caller feeding the final buffered // mel chunk with is_last=true (which keeps the streaming tail frames). At the // session level there is no further audio to process here, so finalize just // returns whatever newly-finalized text remains since the last take. NeMo's // cache-aware streaming behaves identically: the final chunk's incomplete // right context means a trailing is NOT recovered, so we never // fabricate one — finalize emits only the carried-state tokens already // decoded. // // Word side: the end-of-stream has no further word-start markers, so the // trailing open word is now final too — regroup with flush_all so it becomes // available to drain_words(). regroup_words(/*flush_all=*/true); return take_new_text(); } void StreamingSession::regroup_words(bool flush_all, size_t eou_word_tokens) { // Re-run the validated offline grouping over the whole accumulated // non-special token sequence (it does the punctuation lookahead / refinement // exactly like the offline transcribe_with_timestamps path). The last word is // still "open" (its text/end can change when more tokens arrive), so only // words BEFORE it are considered final mid-stream; flush_all makes every word // final at end-of-stream. words_ = group_words(word_tokens_, ml_.config().tokenizer_pieces, frame_sec_f_); if (words_.empty()) { words_finalized_ = 0; } else if (flush_all) { words_finalized_ = words_.size(); } else { words_finalized_ = words_.size() - 1; // An ends the utterance: the decoder restarts from a fresh state // (see feed_mel_chunk), so no later token can extend a word decoded // before it and those words are final now. Without this the trailing // word of every utterance is withheld until the NEXT utterance emits a // token, so each utterance's last word lags by a whole utterance. // Regroup the closed prefix to count them: the next utterance opens on a // word-start token, so its words are exactly words_[closed.size()..]. if (eou_word_tokens > 0) { const std::vector closed = group_words( std::vector(word_tokens_.begin(), word_tokens_.begin() + eou_word_tokens), ml_.config().tokenizer_pieces, frame_sec_f_); if (closed.size() > eou_closed_words_) { eou_closed_words_ = closed.size(); } } // Applied on every chunk, not just the one that carried the : a // caller feeding several chunks per call only drains once at the end, so // recomputing the line above would un-finalize the closed word again and // hold it until the NEXT utterance emits a token. if (eou_closed_words_ > words_finalized_) { words_finalized_ = eou_closed_words_; } } // Never "un-finalize" a word we've already handed out. if (words_finalized_ < words_taken_) words_finalized_ = words_taken_; } std::vector StreamingSession::drain_words() { std::vector out; for (size_t i = words_taken_; i < words_finalized_ && i < words_.size(); ++i) out.push_back(words_[i]); words_taken_ = words_finalized_; return out; } std::string StreamingSession::take_new_text() { if (text_taken_ >= text_.size()) return std::string(); std::string delta = text_.substr(text_taken_); text_taken_ = text_.size(); return delta; } std::vector StreamingSession::drain_events() { std::vector out; out.swap(events_); return out; } void run_stream_over_pcm( StreamingSession& sess, const ModelLoader& ml, const std::vector& pcm16k, const std::function&, const std::vector&)>& on_chunk, const std::string& target_lang) { // target_lang is intentionally unused here: the session already carries its // resolved prompt index from construction. The parameter exists so callers // can route a language through a single entry point (Phase 4 C-API/CLI). (void)target_lang; // 1. Full-clip mel [n_mels, T] (feat-major inner=T), matching the offline / // NeMo online_normalization=False reference (normalization over the whole // clip). The streaming numerics come from the carried encoder/decoder // caches, not from chunking the front end. MelFrontend mel_fe(ml); std::vector mel; int n_mels = 0, T = 0; mel_fe.compute(pcm16k, mel, n_mels, T); if (T <= 0) return; const int chunk0 = sess.chunk_size_first(); // 9 const int chunk_main = sess.chunk_size(); // 16 const int pre_cache = sess.pre_encode_cache_size(); // 9 // mel[:, lo:hi] in feat-major layout. auto window = [&](int lo, int hi) { const int len = hi - lo; std::vector w((size_t)n_mels * len); for (int m = 0; m < n_mels; ++m) for (int t = 0; t < len; ++t) w[(size_t)m * len + t] = mel[(size_t)m * T + (lo + t)]; return w; }; int buffer_idx = 0; bool first = true; while (buffer_idx < T) { const int chunk_size = first ? chunk0 : chunk_main; const int shift = chunk_size; // shift_size == chunk_size here const int chunk_hi = std::min(buffer_idx + chunk_size, T); if (chunk_hi - buffer_idx <= 0) break; const int lo = first ? buffer_idx : std::max(0, buffer_idx - pre_cache); std::vector win = window(lo, chunk_hi); const int win_frames = chunk_hi - lo; const bool is_last = (chunk_hi >= T); sess.feed_mel_chunk(win, win_frames, is_last); if (on_chunk) { std::string nt = sess.take_new_text(); std::vector ev = sess.drain_events(); std::vector wd = sess.drain_words(); if (!nt.empty() || !ev.empty() || !wd.empty()) on_chunk(nt, ev, wd); } buffer_idx += shift; first = false; } } } // namespace pk