#[provable_contracts_macros::contract(
"setfit-encoder-conformance-v1",
equation = "embedding_gather"
)]
pub fn embedding_gather(
weight: &Tensor,
ids: &[u32],
batch: usize,
seq: usize,
) -> Result<Tensor, OpError> {
let w_shape = weight.shape();
if w_shape.len() != 2 {
return Err(OpError::ShapeMismatch {
expected: vec![0, 0],
got: w_shape.to_vec(),
});
}
let vocab_size = w_shape[0];
let hidden = w_shape[1];
if vocab_size == 0 {
return Err(OpError::ZeroDimension {
which: "vocab_size",
});
}
if hidden == 0 {
return Err(OpError::ZeroDimension { which: "hidden" });
}
if batch == 0 {
return Err(OpError::ZeroDimension { which: "batch" });
}
if seq == 0 {
return Err(OpError::ZeroDimension { which: "seq" });
}
let overflow = || OpError::ShapeOverflow {
dims: vec![batch, seq, hidden],
};
let positions = batch.checked_mul(seq).ok_or_else(overflow)?;
let total = positions.checked_mul(hidden).ok_or_else(overflow)?;
if ids.len() != positions {
return Err(OpError::ShapeMismatch {
expected: vec![batch, seq],
got: vec![ids.len()],
});
}
let w = weight.data();
if let Some(position) = w.iter().position(|v| !v.is_finite()) {
return Err(OpError::NonFiniteInput { position });
}
for (position, &id) in ids.iter().enumerate() {
if id as usize >= vocab_size {
return Err(OpError::OutOfVocabulary {
id,
vocab_size,
position,
});
}
}
contract_pre_embedding_gather!(ids);
let mut out = vec![0.0f32; total];
for (i, &id) in ids.iter().enumerate() {
let src = (id as usize) * hidden;
let dst = i * hidden;
out[dst..dst + hidden].copy_from_slice(&w[src..src + hidden]);
}
let mut result = Tensor::from_vec(out, &[batch, seq, hidden]);
if is_grad_enabled() && weight.requires_grad_enabled() {
result.requires_grad_(true);
let grad_fn = Arc::new(EmbeddingBackward {
indices: ids.to_vec(),
vocab_size,
hidden_size: hidden,
});
result.set_grad_fn(grad_fn.clone());
with_graph(|graph| {
graph.register_tensor(weight.clone());
graph.record(result.id(), grad_fn, vec![weight.id()]);
});
}
contract_post_embedding_gather!(result.data());
Ok(result)
}