#[provable_contracts_macros::contract(
"setfit-encoder-conformance-v1",
equation = "masked_mean_pool"
)]
pub fn masked_mean_pool(hidden: &Tensor, mask: &[u8]) -> Result<Tensor, OpError> {
let shape = hidden.shape();
if shape.len() != 3 {
return Err(OpError::ShapeMismatch {
expected: vec![0, 0, 0],
got: shape.to_vec(),
});
}
let batch = shape[0];
let seq = shape[1];
let hidden_size = shape[2];
if batch == 0 {
return Err(OpError::ZeroDimension { which: "batch" });
}
if seq == 0 {
return Err(OpError::ZeroDimension { which: "seq" });
}
if hidden_size == 0 {
return Err(OpError::ZeroDimension { which: "hidden" });
}
let positions = batch.checked_mul(seq).ok_or(OpError::ShapeOverflow {
dims: vec![batch, seq, hidden_size],
})?;
let total = positions
.checked_mul(hidden_size)
.ok_or(OpError::ShapeOverflow {
dims: vec![batch, seq, hidden_size],
})?;
if mask.len() != positions {
return Err(OpError::LengthMismatch {
ids: positions,
mask: mask.len(),
});
}
for (position, &v) in mask.iter().enumerate() {
if v > 1 {
return Err(OpError::NonBinaryMaskValue { value: v, position });
}
}
let mut counts = Vec::with_capacity(batch);
for row in 0..batch {
let base = row * seq;
let count = mask[base..base + seq]
.iter()
.fold(0usize, |acc, &m| acc + usize::from(m == 1));
if count == 0 {
return Err(OpError::AllPaddingRow { row });
}
counts.push(count);
}
contract_pre_masked_mean_pool!(mask);
debug_assert_eq!(hidden.numel(), total, "shape product must match numel");
let x = hidden.data();
let mut out = vec![0.0f32; batch * hidden_size];
for (row, &count) in counts.iter().enumerate() {
let base = row * seq;
let out_off = row * hidden_size;
for pos in 0..seq {
if mask[base + pos] != 1 {
continue;
}
let src = base * hidden_size + pos * hidden_size;
for j in 0..hidden_size {
out[out_off + j] += x[src + j];
}
}
let inv = 1.0 / count as f32;
for j in 0..hidden_size {
out[out_off + j] *= inv;
}
}
let mut result = Tensor::from_vec(out, &[batch, hidden_size]);
if is_grad_enabled() && hidden.requires_grad_enabled() {
result.requires_grad_(true);
let grad_fn = Arc::new(MaskedMeanPoolBackward {
mask: mask.to_vec(),
batch,
seq,
hidden: hidden_size,
});
result.set_grad_fn(grad_fn.clone());
with_graph(|graph| {
graph.register_tensor(hidden.clone());
graph.record(result.id(), grad_fn, vec![hidden.id()]);
});
}
contract_post_masked_mean_pool!(result.data());
Ok(result)
}