diff --git a/src/models/vlm/sam3_image/impl.rs b/src/models/vlm/sam3_image/impl.rs index 6881a192..5f42bd89 100644 --- a/src/models/vlm/sam3_image/impl.rs +++ b/src/models/vlm/sam3_image/impl.rs @@ -348,12 +348,30 @@ impl Sam3Image { let mut uncached = Vec::new(); let mut seen = HashSet::new(); for p in prompts { - if seen.insert(&p.text) && !self.text_cache.contains(&p.text) { + if seen.insert(&p.text) && self.text_cache.get(&p.text).is_none() { uncached.push(p.text.as_str()); } } + + let unique_prompts_count = seen.len(); + if unique_prompts_count > self.text_cache.cap().get() { + // If the number of unique prompts in the current batch exceeds the cache capacity, force an expansion. + self.text_cache + .resize(NonZeroUsize::new(unique_prompts_count).unwrap()); + } + if !uncached.is_empty() { - for (text, feat) in uncached.iter().zip(self.encode_texts(engines, &uncached)?) { + let encoded_feats = self.encode_texts(engines, &uncached)?; + + if encoded_feats.len() != uncached.len() { + bail!( + "Encoded features length mismatch: expected {}, got {}", + uncached.len(), + encoded_feats.len() + ); + } + + for (text, feat) in uncached.iter().zip(encoded_feats) { self.text_cache.put(text.to_string(), feat); } }