Repository navigation
[PocketTTS] Add seed support and voice embedding caching for consiste… - #3189
Conversation
…nt generation Introduce a thread-safe LRU cache for voice embeddings to skip redundant Mimi encoder runs: compute a sampled hash of reference audio, return cached Ort tensor on hit, and store embeddings on miss. Expose voice_embedding_cache_capacity through model config, C API and Go wrapper, and add CLI flag pocket-voice-embedding-cache-capacity. Add deterministic NormalDataGenerator constructor (seed-aware) and use the generation 'seed' option to produce reproducible noise; preserve original thread-local RNG when seed < 0. Update config ToString and defaults, add necessary headers and logging/timing for cache hits/misses.
|
Note Reviews pausedIt looks like this branch is under active development. To avoid overwhelming you with review comments due to an influx of new commits, CodeRabbit has automatically paused this review. You can configure this behavior by changing the Use the following commands to manage reviews:
Use the checkboxes below for quick actions:
📝 WalkthroughWalkthroughAdds a voice-embedding LRU cache to Offline TTS with a configurable capacity exposed through C/C++/Go APIs, and adds deterministic seed support to NormalDataGenerator for repeatable generation. Changes
Sequence Diagram(s)sequenceDiagram
autonumber
participant Client
participant OfflineTTS as OfflineTtsPocketImpl
participant Cache as VoiceEmbeddingCache
participant Encoder as MimiEncoder
participant NormalGen as NormalDataGenerator
Client->>OfflineTTS: synthesize(ref_audio, gen_config{seed})
OfflineTTS->>Cache: compute audio_hash -> Get(hash)
alt cache hit
Cache-->>OfflineTTS: embedding
OfflineTTS->>NormalGen: generate noise (seed forwarded)
OfflineTTS-->>Client: return synthesized audio (uses cached embedding)
else cache miss
OfflineTTS->>Encoder: Run encoder on ref_audio
Encoder-->>OfflineTTS: embedding
OfflineTTS->>Cache: Put(hash, embedding)
OfflineTTS->>NormalGen: generate noise (seed forwarded)
OfflineTTS-->>Client: return synthesized audio
end
Estimated code review effort🎯 4 (Complex) | ⏱️ ~45 minutes Possibly related issues
Possibly related PRs
Suggested reviewers
Poem
🚥 Pre-merge checks | ✅ 2 | ❌ 1❌ Failed checks (1 warning)
✅ Passed checks (2 passed)
✏️ Tip: You can configure your own custom pre-merge checks in the settings. ✨ Finishing Touches🧪 Generate unit tests (beta)
Thanks for using CodeRabbit! It's free for OSS, and your support helps us grow. If you like it, consider giving us a shout-out. Comment |
Summary of ChangesHello @ramishi, I'm Gemini Code Assist1! I'm currently reviewing this pull request and will post my feedback shortly. In the meantime, here's a summary to help you and other reviewers quickly get up to speed! This pull request significantly enhances the PocketTTS system by introducing a robust caching mechanism for voice embeddings and enabling deterministic noise generation. The caching feature aims to boost performance by avoiding repetitive computations, while seed support ensures reproducibility in noise generation. These changes collectively improve the efficiency and reliability of the TTS system, making it more predictable and faster for repeated voice embedding requests. Highlights
Changelog
Activity
Using Gemini Code AssistThe full guide for Gemini Code Assist can be found on our documentation page, here are some quick tips. Invoking Gemini You can request assistance from Gemini at any point by creating a comment using either
Customization To customize Gemini Code Assist for GitHub experience, repository maintainers can create a configuration file and/or provide a custom code review style guide (such as PEP-8 for Python) by creating and adding files to a Limitations & Feedback Gemini Code Assist may make mistakes. Please leave feedback on any instances where its feedback is incorrect or counter productive. You can react with 👍 and 👎 on @gemini-code-assist comments. If you're interested in giving your feedback about your experience with Gemini Code Assist for Github and other Google products, sign up here. You can also get AI-powered code generation, chat, as well as code reviews directly in the IDE at no cost with the Gemini Code Assist IDE Extension. Footnotes
|
|
from issue: #3185 |
There was a problem hiding this comment.
Code Review
This pull request introduces two significant enhancements to PocketTTS: support for seeded random number generation for reproducible outputs, and a thread-safe LRU cache for voice embeddings to optimize performance by avoiding repeated computations. The changes are well-integrated across the C++ core, C API, and Go bindings. My review includes a couple of minor suggestions to improve code style and efficiency in the C++ implementation of the caching logic. Overall, this is a solid contribution that improves both functionality and performance.
| size_t total = 1; | ||
| for (auto s : result_shape) total *= s; |
There was a problem hiding this comment.
| void Put(size_t key, const std::vector<float> &in_data, | ||
| const std::vector<int64_t> &in_shape) { | ||
| std::lock_guard<std::mutex> lock(mutex); | ||
| if (capacity == 0) { | ||
| return; | ||
| } | ||
|
|
||
| auto it = data.find(key); | ||
| if (it != data.end()) { | ||
| // Update existing entry | ||
| it->second = {in_data, in_shape}; | ||
| lru_list.splice(lru_list.begin(), lru_list, map_iters[key]); | ||
| return; | ||
| } | ||
|
|
||
| // Evict if necessary | ||
| if (data.size() >= capacity && !lru_list.empty()) { | ||
| size_t last = lru_list.back(); | ||
| data.erase(last); | ||
| map_iters.erase(last); | ||
| lru_list.pop_back(); | ||
| } | ||
|
|
||
| // Insert new entry | ||
| lru_list.push_front(key); | ||
| data[key] = {in_data, in_shape}; | ||
| map_iters[key] = lru_list.begin(); | ||
| } |
There was a problem hiding this comment.
To improve efficiency, consider passing in_data and in_shape by value to the Put method and then using std::move to transfer ownership. This can avoid unnecessary copies, especially since the call site creates temporary std::vector objects.
void Put(size_t key, std::vector<float> in_data,
std::vector<int64_t> in_shape) {
std::lock_guard<std::mutex> lock(mutex);
if (capacity == 0) {
return;
}
auto it = data.find(key);
if (it != data.end()) {
// Update existing entry
it->second = {std::move(in_data), std::move(in_shape)};
lru_list.splice(lru_list.begin(), lru_list, map_iters[key]);
return;
}
// Evict if necessary
if (data.size() >= capacity && !lru_list.empty()) {
size_t last = lru_list.back();
data.erase(last);
map_iters.erase(last);
lru_list.pop_back();
}
// Insert new entry
lru_list.push_front(key);
data[key] = {std::move(in_data), std::move(in_shape)};
map_iters[key] = lru_list.begin();
}There was a problem hiding this comment.
Actionable comments posted: 5
Caution
Some comments are outside the diff and can’t be posted inline due to platform limitations.
⚠️ Outside diff range comments (1)
sherpa-onnx/csrc/offline-tts-pocket-impl.h (1)
95-101:⚠️ Potential issue | 🟠 MajorCache capacity is only set in the Manager constructor.
The standard constructor never calls
GetCache().SetCapacity(...), so the configured capacity is ignored in the common path.🔧 Suggested fix
explicit OfflineTtsPocketImpl(const OfflineTtsConfig &config) : config_(config), model_(std::make_unique<OfflineTtsPocketModel>(config.model)) { InitTokenizer(); + GetCache().SetCapacity(config.model.pocket.voice_embedding_cache_capacity);
🤖 Fix all issues with AI agents
In `@sherpa-onnx/c-api/c-api.cc`:
- Around line 1366-1367: The current use of SHERPA_ONNX_OR when setting
tts_config.model.pocket.voice_embedding_cache_capacity forces 0 to be treated as
"unset" so callers can't disable caching; change the assignment to preserve an
explicit 0 by using an explicit presence check instead of SHERPA_ONNX_OR — e.g.
if the config struct exposes a has_voice_embedding_cache_capacity flag use that
to decide between config->model.pocket.voice_embedding_cache_capacity and 50, or
if the codebase uses a sentinel (e.g. -1) treat -1 as "unset" and otherwise
assign the config value (allowing 0 through). Ensure you update the assignment
referencing tts_config.model.pocket.voice_embedding_cache_capacity and
config->model.pocket.voice_embedding_cache_capacity accordingly.
In `@sherpa-onnx/csrc/normal-data-generator.cc`:
- Around line 37-55: The Fill method currently constructs a local std::mt19937
each call when seed_ >= 0 which produces identical sequences; instead add and
use an instance-level RNG (e.g. a member std::mt19937 rng_ initialized from
seed_ in the NormalDataGenerator constructor when seed_ >= 0) and use it in
NormalDataGenerator::Fill rather than creating a local rng; you can keep the
existing RNGHolder-based thread-local path for the seed_ < 0 case and continue
updating the distribution parameters (mean_, stddev_) before sampling from the
appropriate RNG.
In `@sherpa-onnx/csrc/offline-tts-pocket-impl.h`:
- Around line 544-549: The current cache key uses a hash of
gen_config.reference_audio pre-processing and omits reference_sample_rate and
effective sample count, causing collisions across different resampling/trimming;
fix by computing the hash after any resampling/truncation (i.e., on the
post-processed buffer used to compute embeddings) and mix in
reference_sample_rate and the final number of samples (or
max_reference_audio_len/num_samples) into the hash; update the hashing code that
builds audio_hash (used in the embedding cache lookup) to incorporate these
values (e.g., combine std::hash<float> over the post-processed samples and then
xor/mix in std::hash<size_t>(reference_sample_rate) and
std::hash<size_t>(num_samples) or effective length).
- Around line 553-572: The code on cache hit constructs data_copy and shape_copy
as local buffers then calls Ort::Value::CreateTensor with that external pointer,
causing a use-after-free when the function returns; instead allocate an
Ort-owned tensor and copy the cached bytes into its internal buffer before
returning. Concretely, replace the CreateTensor call that uses data_copy.data()
with creating a tensor that owns its memory (use the Ort allocator or
CreateTensor with allocated buffer), obtain the tensor's mutable data pointer
(e.g., GetTensorMutableData) and memcpy or std::copy from data_copy into that
buffer, then return the owning Ort::Value; update the cache-hit path around
GetCache().Get, data_copy, shape_copy, memory_info, and result accordingly.
In `@sherpa-onnx/csrc/offline-tts-pocket-model-config.cc`:
- Around line 26-29: Validate voice_embedding_cache_capacity to prevent
negative-to-size_t wrap: after options are parsed (or immediately before
converting to size_t where the cache is constructed), check that
voice_embedding_cache_capacity >= 0; if it's negative, either clamp it to 0
(treat 0 as "disable caching") or return an error/log and exit. Update the
Register/help text for voice_embedding_cache_capacity to mention "Default: 50; 0
disables caching" and ensure any place that does
static_cast<size_t>(voice_embedding_cache_capacity) only runs after this
validation.
Several fixes and improvements across offline TTS and utilities: - c-api: minor formatting cleanup when copying segment timestamps; handle pocket voice_embedding_cache_capacity explicitly (use provided non-negative value or default to 50). - normal-data-generator: switch to an instance-level std::mt19937 seeded in the constructor for deterministic mode and use it in Fill(); added <random> and rng_ member. - offline-tts-pocket-impl.h: strengthen cache key by mixing in sample rate and length to reduce collisions; create an owned Ort tensor via AllocatorWithDefaultOptions and CreateTensor<float>(), then copy cached data into the tensor to avoid potential use-after-free. - offline-tts-pocket-model-config.cc: update help text to note that 0 disables caching and add validation to reject negative voice_embedding_cache_capacity with an error log. These changes improve determinism, correctness of cached tensor lifetimes, and robustness of the cache key and configuration validation.
There was a problem hiding this comment.
Actionable comments posted: 1
Caution
Some comments are outside the diff and can’t be posted inline due to platform limitations.
⚠️ Outside diff range comments (1)
sherpa-onnx/csrc/offline-tts-pocket-impl.h (1)
40-43:⚠️ Potential issue | 🟠 Major
SetCapacitynot called in the non-manager constructor.The manager constructor (line 100) calls
GetCache().SetCapacity(...), but this constructor does not. The cache will use its hardcoded default (50) regardless ofconfig.model.pocket.voice_embedding_cache_capacity.Proposed fix
explicit OfflineTtsPocketImpl(const OfflineTtsConfig &config) : config_(config), model_(std::make_unique<OfflineTtsPocketModel>(config.model)) { InitTokenizer(); + GetCache().SetCapacity(config.model.pocket.voice_embedding_cache_capacity);
🤖 Fix all issues with AI agents
In `@sherpa-onnx/csrc/offline-tts-pocket-impl.h`:
- Around line 544-555: The cache key currently omits
gen_config.max_reference_audio_len so different truncation lengths can produce
identical audio_hash; update the hash computation that produces audio_hash to
incorporate gen_config.max_reference_audio_len (or compute the hash after
truncating reference_audio to max_reference_audio_len) so the key reflects the
actual audio used for embedding; locate the audio_hash construction near the
reference_audio sampling loop (symbols: audio_hash, gen_config.reference_audio,
gen_config.reference_sample_rate) and mix in
std::hash<size_t>{}(gen_config.max_reference_audio_len) (or use truncated
buffer) before returning/using audio_hash.
🧹 Nitpick comments (3)
sherpa-onnx/csrc/offline-tts-pocket-impl.h (3)
574-577:%lldformat specifier may trigger warnings on some platforms.
std::chrono::microseconds::repis not guaranteed to belong long. On some targets (e.g., 32-bit or MSVC with certain configs), this can produce format-string warnings. Consider usingstatic_cast<long long>(...)consistently orPRId64if targetingint64_t.This also applies to the
%lldon line 651.
557-579: Cache hit/miss logging usesSHERPA_ONNX_LOGE(error level) for informational messages.Cache hit/miss events are normal operational behavior, not errors. Using the error-level logger for these will pollute error logs in production. If a debug/info-level log macro is available in this codebase, prefer that (or gate behind
config_.model.debug).Also applies to: 634-654
827-830: Static cache is shared across allOfflineTtsPocketImplinstances — deliberate?
GetCache()returns astaticlocal, so all instances share one cache. If two instances use different models or configs, they could serve each other's cached embeddings. This is fine if there's only ever one active model, but worth documenting the assumption.
Replace manual product loop with std::accumulate + std::multiplies for computing tensor element count. Change Cache::Put to take vectors by value and std::move them into the map (both on update and insert) to avoid unnecessary copies and improve performance.
There was a problem hiding this comment.
Actionable comments posted: 3
Caution
Some comments are outside the diff and can’t be posted inline due to platform limitations.
⚠️ Outside diff range comments (1)
sherpa-onnx/csrc/offline-tts-pocket-impl.h (1)
40-93:⚠️ Potential issue | 🟠 MajorCache capacity is never set in the non-Manager constructor.
The Manager-template constructor (line 100) calls
GetCache().SetCapacity(config.model.pocket.voice_embedding_cache_capacity), but this constructor does not. Users going through this code path will always get the hardcoded default capacity (50), ignoring their configuration.Proposed fix
InitTokenizer(); + GetCache().SetCapacity(config.model.pocket.voice_embedding_cache_capacity); if (!config.rule_fsts.empty()) {
🤖 Fix all issues with AI agents
In `@sherpa-onnx/csrc/offline-tts-pocket-impl.h`:
- Around line 574-577: The cache hit/miss messages currently use the error-level
macro SHERPA_ONNX_LOGE; change these to a non-error log (e.g.,
SHERPA_ONNX_LOGD/LOGI if available) or wrap the SHERPA_ONNX_LOGE calls with the
same config_.model.debug guard used elsewhere so they only log when debug is
enabled (apply the same fix for the other occurrence referenced near the Mimi
encoder cache logging). Locate the cache logging calls that print "CACHE HIT:
voice embedding..." and the corresponding miss log and replace the macro or add
a conditional check using config_.model.debug to avoid emitting error-level logs
for normal cache events.
- Around line 827-830: GetCache() currently returns a process-wide
VoiceEmbeddingCache singleton which can be contaminated across different
OfflineTtsPocketImpl instances and model types; change this by making the cache
model-aware or instance-scoped: either (a) include a unique model identifier
(e.g., Mimi encoder model name or path) into the cache key construction so
GetCache() still returns a shared cache but keys are namespaced per model, or
(b) remove the static singleton and add a VoiceEmbeddingCache member (or
std::shared_ptr<VoiceEmbeddingCache>) to OfflineTtsPocketImpl so each instance
controls capacity via SetCapacity without affecting others; update all places
that call GetCache() and any cache key generation to incorporate the chosen
model id or to use the instance member instead.
- Around line 552-555: The current sampling loop that updates audio_hash by
hashing every 100th float (iterating over gen_config.reference_audio with step
100) is a weak cache key causing silent collisions; change it to compute the
hash over the entire reference buffer (iterate over all
gen_config.reference_audio elements) or replace the sampling scheme with a fast
bulk checksum/hash (e.g., a vector- or byte-oriented hash) to ensure unique keys
for gen_config.reference_audio before using audio_hash as the embedding cache
key.
🧹 Nitpick comments (1)
sherpa-onnx/csrc/offline-tts-pocket-impl.h (1)
546-550:reference_audio.size()is mixed into the hash twice.Line 546 initializes
audio_hashto.size(), and lines 549–550 hash.size()again viastd::hash<int32_t>. The second mix adds no entropy beyond what the seed already provides. Consider initializing with a fixed seed (e.g.,0or a prime) and hashing the size once, which is clearer.
- Robustify voice embedding cache key: include `max_reference_audio_len` and hash all reference audio samples to prevent collisions. - Optimization: Switch cache hit/miss logs from ERROR to INFO and make timing logic conditional on debug mode. - Refactor: Merge [NormalDataGenerator] constructors and remove unused `<strstream>` header.
There was a problem hiding this comment.
🧹 Nitpick comments (4)
sherpa-onnx/csrc/normal-data-generator.cc (1)
37-54: Thread-safety note:mutable rng_in aconstmethod is not thread-safe for the deterministic path.If the same
NormalDataGeneratorinstance (withseed_ >= 0) is ever called from multiple threads, there's a data race onrng_. This is fine given current usage (local variable inGenerateSingleSentence), but worth noting if the usage pattern ever changes.sherpa-onnx/csrc/offline-tts-pocket-impl.h (3)
548-552: Redundant hashing ofreference_audio.size().
audio_hashis initialized withreference_audio.size()on line 548, thenreference_audio.size()is hashed again (asint32_t) on lines 551-552. The second hash is redundant. Also note the implicit narrowing fromsize_ttoint32_ton line 551.Proposed cleanup
size_t audio_hash = gen_config.reference_audio.size(); audio_hash ^= std::hash<int32_t>{}(gen_config.reference_sample_rate) + 0x9e3779b9 + (audio_hash << 6) + (audio_hash >> 2); - audio_hash ^= std::hash<int32_t>{}(gen_config.reference_audio.size()) + - 0x9e3779b9 + (audio_hash << 6) + (audio_hash >> 2);
566-598: Reduce duplication between debug and non-debug cache-hit branches.Lines 571-588 and 589-597 differ only by the timing/logging. The tensor construction and data copy are identical. Consider extracting the common logic.
Proposed refactor
if (cache_.Get(audio_hash, data_copy, shape_copy)) { + // Create an owned tensor and copy data to avoid use-after-free + Ort::AllocatorWithDefaultOptions allocator; + auto result = Ort::Value::CreateTensor<float>( + allocator, shape_copy.data(), shape_copy.size()); + std::copy(data_copy.begin(), data_copy.end(), + result.GetTensorMutableData<float>()); if (config_.model.debug) { - auto cache_start = std::chrono::high_resolution_clock::now(); - - // Create an owned tensor and copy data to avoid use-after-free - Ort::AllocatorWithDefaultOptions allocator; - auto result = Ort::Value::CreateTensor<float>( - allocator, shape_copy.data(), shape_copy.size()); - std::copy(data_copy.begin(), data_copy.end(), - result.GetTensorMutableData<float>()); - auto cache_us = - std::chrono::duration_cast<std::chrono::microseconds>( - std::chrono::high_resolution_clock::now() - cache_start) - .count(); - SHERPA_ONNX_LOG(INFO) - << "CACHE HIT: voice embedding (hash=" << audio_hash - << ") returned in " << (long long)cache_us - << " us (Mimi encoder SKIPPED)"; - return result; - } else { - // Create an owned tensor and copy data to avoid use-after-free - Ort::AllocatorWithDefaultOptions allocator; - auto result = Ort::Value::CreateTensor<float>( - allocator, shape_copy.data(), shape_copy.size()); - std::copy(data_copy.begin(), data_copy.end(), - result.GetTensorMutableData<float>()); - return result; + SHERPA_ONNX_LOG(INFO) + << "CACHE HIT: voice embedding (hash=" << audio_hash + << ") (Mimi encoder SKIPPED)"; } + return result; }Note: The current debug-path timing only measures tensor creation/copy, not the cache lookup itself—so the reported microseconds aren't very meaningful. If cache-hit timing is needed, start the timer before
cache_.Get().
663-693: Same duplication pattern on the cache-miss path.The debug and non-debug branches (lines 663-693) repeat the same caching logic. Consider refactoring similarly:
Proposed refactor
- if (config_.model.debug) { - auto embed_ms = - std::chrono::duration_cast<std::chrono::milliseconds>( - std::chrono::high_resolution_clock::now() - embed_start) - .count(); - - // Cache the embedding in shared LRU cache - auto result_shape = result.GetTensorTypeAndShapeInfo().GetShape(); - size_t total = std::accumulate(result_shape.begin(), result_shape.end(), - 1, std::multiplies<size_t>()); - const float *result_data = result.GetTensorData<float>(); - - cache_.Put( - audio_hash, std::vector<float>(result_data, result_data + total), - std::vector<int64_t>(result_shape.begin(), result_shape.end())); - - SHERPA_ONNX_LOG(INFO) - << "CACHE MISS: Mimi encoder took " << (long long)embed_ms - << " ms → cached embedding (hash=" << audio_hash << ", " << total - << " floats)"; - } else { - // Cache the embedding in shared LRU cache (no logging) - auto result_shape = result.GetTensorTypeAndShapeInfo().GetShape(); - size_t total = std::accumulate(result_shape.begin(), result_shape.end(), - 1, std::multiplies<size_t>()); - const float *result_data = result.GetTensorData<float>(); - - cache_.Put( - audio_hash, std::vector<float>(result_data, result_data + total), - std::vector<int64_t>(result_shape.begin(), result_shape.end())); - } + auto result_shape = result.GetTensorTypeAndShapeInfo().GetShape(); + size_t total = std::accumulate(result_shape.begin(), result_shape.end(), + 1, std::multiplies<size_t>()); + const float *result_data = result.GetTensorData<float>(); + + cache_.Put( + audio_hash, std::vector<float>(result_data, result_data + total), + std::vector<int64_t>(result_shape.begin(), result_shape.end())); + + if (config_.model.debug) { + auto embed_ms = + std::chrono::duration_cast<std::chrono::milliseconds>( + std::chrono::high_resolution_clock::now() - embed_start) + .count(); + SHERPA_ONNX_LOG(INFO) + << "CACHE MISS: Mimi encoder took " << (long long)embed_ms + << " ms → cached embedding (hash=" << audio_hash << ", " << total + << " floats)"; + }
Co-authored-by: Fangjun Kuang <csukuangfj@gmail.com>
There was a problem hiding this comment.
Actionable comments posted: 1
🧹 Nitpick comments (1)
sherpa-onnx/csrc/normal-data-generator.cc (1)
40-43: Recreatingstd::normal_distributionon everyFill()call discards cached variate
std::normal_distributionimplementations commonly use Box-Muller (generating two values, caching one). Constructing a freshdiston eachFill()call discards any cached value and forces the RNG to advance an extra step, which can affect bit-for-bit reproducibility across different call patterns.Since
rng_is already amutableinstance member, adding amutable std::normal_distribution<float> dist_alongside it (initialized in the constructor) would be fully consistent with the rest of the fix and eliminate the overhead:♻️ Proposed refactor
In
sherpa-onnx/csrc/normal-data-generator.h, add adist_member alongsiderng_:mutable std::mt19937 rng_; // used if seed_ >= 0 + mutable std::normal_distribution<float> dist_; // used if seed_ >= 0In
sherpa-onnx/csrc/normal-data-generator.cc, initialize it in the constructor and reuse it inFill():NormalDataGenerator::NormalDataGenerator(float mean /*= 0.0f*/, float stddev /*= 1.0f*/, int32_t seed /*= -1*/) - : mean_(mean), stddev_(stddev), seed_(seed) { + : mean_(mean), stddev_(stddev), seed_(seed), dist_(mean, stddev) { if (seed_ >= 0) { rng_.seed(static_cast<unsigned>(seed_)); } } void NormalDataGenerator::Fill(float *data, std::size_t size) const { if (seed_ >= 0) { - std::normal_distribution<float> dist(mean_, stddev_); for (std::size_t i = 0; i < size; ++i) { - data[i] = dist(rng_); + data[i] = dist_(rng_); }🤖 Prompt for AI Agents
Verify each finding against the current code and only fix it if needed. In `@sherpa-onnx/csrc/normal-data-generator.cc` around lines 40 - 43, The Fill() function currently constructs a local std::normal_distribution<float> (dist) each call which discards the distribution's cached variate; instead add a mutable std::normal_distribution<float> dist_ as a member (alongside rng_), declare it in normal-data-generator.h, initialize dist_ with mean_ and stddev_ in the class constructor's initializer list in normal-data-generator.cc, and modify Fill() to use the member dist_ rather than creating a new local distribution so cached variates are preserved and RNG advances remain consistent.
🤖 Prompt for all review comments with AI agents
Verify each finding against the current code and only fix it if needed.
Inline comments:
In `@sherpa-onnx/csrc/normal-data-generator.cc`:
- Around line 30-35: The constructor definition
NormalDataGenerator::NormalDataGenerator(float mean = /*0.0f*/, float stddev =
/*1.0f*/, int32_t seed = /*-1*/) must not contain default-argument initializers;
remove the `= /*...*/` pieces so the signature is simply
NormalDataGenerator::NormalDataGenerator(float mean, float stddev, int32_t seed)
and keep the defaults only in the class declaration in the header; if you want
inline documentation of the defaults, use comment-before-equals style like
/*=0.0f*/ next to the parameter names in the definition, and leave the body that
sets mean_, stddev_, seed_ and seeds rng_ unchanged.
---
Nitpick comments:
In `@sherpa-onnx/csrc/normal-data-generator.cc`:
- Around line 40-43: The Fill() function currently constructs a local
std::normal_distribution<float> (dist) each call which discards the
distribution's cached variate; instead add a mutable
std::normal_distribution<float> dist_ as a member (alongside rng_), declare it
in normal-data-generator.h, initialize dist_ with mean_ and stddev_ in the class
constructor's initializer list in normal-data-generator.cc, and modify Fill() to
use the member dist_ rather than creating a new local distribution so cached
variates are preserved and RNG advances remain consistent.
csukuangfj
left a comment
There was a problem hiding this comment.
Thanks! Left some minor comments. Otherwise, it looks great to me.
Co-authored-by: Fangjun Kuang <csukuangfj@gmail.com>
Co-authored-by: Fangjun Kuang <csukuangfj@gmail.com>
|
Thanks for your helps, I've learnt a lot! |
csukuangfj
left a comment
There was a problem hiding this comment.
Thank you for your contribution!
Introduce a thread-safe LRU cache for voice embeddings to skip redundant Mimi encoder runs: compute a sampled hash of reference audio, return cached tensor on hit, and store embeddings on miss. Expose voice_embedding_cache_capacity through model config, C API and Go wrapper, and add CLI flag pocket-voice-embedding-cache-capacity. Add deterministic NormalDataGenerator constructor (seed-aware) and use the generation 'seed' option to produce reproducible noise; preserve original thread-local RNG when seed < 0. Update config ToString and defaults, add necessary headers and logging/timing for cache hits/misses.
Summary by CodeRabbit