use std::{cmp::max, ops::Range};
use futures::{StreamExt, stream};
use crate::driver::DynModel;
use crate::error::ProviderError;
use crate::operation::Embedding as EmbeddingOp;
use crate::{
completion::Usage,
embeddings::{Embed, EmbedError, Embedding, EmbeddingResponse, embed::TextEmbedder},
};
#[must_use = "an embeddings builder does nothing until built"]
pub struct EmbeddingsBuilder<T> {
model: DynModel<EmbeddingOp>,
documents: Vec<(T, Vec<String>)>,
}
impl<T: Embed> EmbeddingsBuilder<T> {
pub fn new(model: impl Into<DynModel<EmbeddingOp>>) -> Self {
Self {
model: model.into(),
documents: vec![],
}
}
pub fn document(mut self, document: T) -> Result<Self, EmbedError> {
let mut embedder = TextEmbedder::default();
document.embed(&mut embedder)?;
self.documents.push((document, embedder.texts));
Ok(self)
}
pub fn documents(self, documents: impl IntoIterator<Item = T>) -> Result<Self, EmbedError> {
let builder = documents
.into_iter()
.try_fold(self, EmbeddingsBuilder::document)?;
Ok(builder)
}
}
impl<T> EmbeddingsBuilder<T>
where
T: Embed + crate::wasm_compat::WasmCompatSend,
{
pub async fn build(self) -> Result<Vec<(T, Vec<Embedding>)>, ProviderError> {
let (result, _usage) = self.build_with_usage().await?;
Ok(result)
}
pub(crate) async fn build_with_usage(
self,
) -> Result<(Vec<(T, Vec<Embedding>)>, Usage), ProviderError> {
use stream::TryStreamExt;
let mut docs: Vec<T> = Vec::with_capacity(self.documents.len());
let mut spans: Vec<Range<usize>> = Vec::with_capacity(self.documents.len());
let mut texts: Vec<String> = Vec::new();
for (doc, doc_texts) in self.documents {
let start = texts.len();
texts.extend(doc_texts);
spans.push(start..texts.len());
docs.push(doc);
}
let total_texts = texts.len();
let max_documents = max(1, self.model.capabilities().max_documents);
let (slots, usage) = stream::iter(texts.into_iter().enumerate())
.chunks(max_documents)
.map(|chunk| async {
let (slots, batch): (Vec<usize>, Vec<String>) = chunk.into_iter().unzip();
let response: EmbeddingResponse = self.model.call(batch).await?;
Ok::<_, ProviderError>((
slots
.into_iter()
.zip(response.embeddings)
.collect::<Vec<_>>(),
response.usage,
))
})
.buffer_unordered(max(1, 1024 / max_documents))
.try_fold(
(
(0..total_texts)
.map(|_| None)
.collect::<Vec<Option<Embedding>>>(),
Usage::default(),
),
|(mut slots, mut usage_acc), (chunk_embeddings, chunk_usage)| async move {
for (slot, embedding) in chunk_embeddings {
if let Some(place) = slots.get_mut(slot) {
*place = Some(embedding);
}
}
usage_acc += chunk_usage;
Ok((slots, usage_acc))
},
)
.await?;
let mut slots = slots.into_iter();
let mut result = Vec::with_capacity(docs.len());
for (index, (doc, span)) in docs.into_iter().zip(spans).enumerate() {
if span.is_empty() {
return Err(crate::error::ProviderError::Response(format!(
"document {index} produced no text to embed, so it has no \
embeddings to return; an empty collection in an `#[embed]` \
field embeds nothing"
)));
}
let embeddings = slots
.by_ref()
.take(span.len())
.collect::<Option<Vec<Embedding>>>()
.ok_or_else(|| {
crate::error::ProviderError::Response(format!(
"provider returned fewer embeddings than texts sent: \
document {index} is missing at least one of its {} texts \
(slots {}..{} of {total_texts})",
span.len(),
span.start,
span.end
))
})?;
result.push((doc, embeddings));
}
Ok((result, usage))
}
}
#[cfg(test)]
mod tests;