#[provable_contracts_macros::contract(
"setfit-encoder-conformance-v1",
equation = "cosine_similarity_rows"
)]
pub fn cosine_similarity_rows(a: &Tensor, b: &Tensor, eps: f32) -> Result<Tensor, OpError> {
if a.shape().len() != 2 {
return Err(OpError::ShapeMismatch {
expected: vec![0, 0],
got: a.shape().to_vec(),
});
}
if b.shape().len() != 2 {
return Err(OpError::ShapeMismatch {
expected: vec![0, 0],
got: b.shape().to_vec(),
});
}
if a.shape() != b.shape() {
return Err(OpError::ShapeMismatch {
expected: a.shape().to_vec(),
got: b.shape().to_vec(),
});
}
let batch = a.shape()[0];
let hidden = a.shape()[1];
if batch == 0 {
return Err(OpError::ZeroDimension { which: "batch" });
}
if hidden == 0 {
return Err(OpError::ZeroDimension { which: "hidden" });
}
if !(eps.is_finite() && eps > 0.0) {
return Err(OpError::invalid_epsilon(eps));
}
let ad = a.data();
let bd = b.data();
if let Some(position) = ad.iter().position(|v| !v.is_finite()) {
return Err(OpError::NonFiniteInput { position });
}
if let Some(position) = bd.iter().position(|v| !v.is_finite()) {
return Err(OpError::NonFiniteInput { position });
}
contract_pre_cosine_similarity_rows!(ad);
let mut out = vec![0.0f32; batch];
let mut norms_a = Vec::with_capacity(batch);
let mut norms_b = Vec::with_capacity(batch);
for row in 0..batch {
let base = row * hidden;
let (mut sa, mut sb, mut dot) = (0.0f64, 0.0f64, 0.0f64);
for j in 0..hidden {
let (x, y) = (f64::from(ad[base + j]), f64::from(bd[base + j]));
sa += x * x;
sb += y * y;
dot += x * y;
}
let na = sa.sqrt() as f32;
let nb = sb.sqrt() as f32;
let da = if na > eps { na } else { eps };
let db = if nb > eps { nb } else { eps };
out[row] = (dot / (f64::from(da) * f64::from(db))) as f32;
norms_a.push(na);
norms_b.push(nb);
}
let mut result = Tensor::from_vec(out, &[batch]);
if is_grad_enabled() && (a.requires_grad_enabled() || b.requires_grad_enabled()) {
result.requires_grad_(true);
let grad_fn = Arc::new(CosineSimilarityBackward {
a: a.clone(),
b: b.clone(),
similarity: result.clone(),
norms_a,
norms_b,
eps,
batch,
hidden,
});
result.set_grad_fn(grad_fn.clone());
with_graph(|graph| {
graph.register_tensor(a.clone());
graph.register_tensor(b.clone());
graph.record(result.id(), grad_fn, vec![a.id(), b.id()]);
});
}
contract_post_cosine_similarity_rows!(result.data());
Ok(result)
}
#[provable_contracts_macros::contract("setfit-encoder-conformance-v1", equation = "mse_loss")]
pub fn mse_loss(pred: &Tensor, target: &[f32]) -> Result<Tensor, OpError> {
if pred.shape().len() != 1 {
return Err(OpError::ShapeMismatch {
expected: vec![0],
got: pred.shape().to_vec(),
});
}
let n = pred.shape()[0];
if n == 0 {
return Err(OpError::ZeroDimension { which: "batch" });
}
if target.len() != n {
return Err(OpError::LengthMismatch {
ids: n,
mask: target.len(),
});
}
if let Some(position) = target.iter().position(|v| !v.is_finite()) {
return Err(OpError::NonFiniteInput { position });
}
contract_pre_mse_loss!(target);
let p = pred.data();
let inv_n = 1.0f64 / n as f64;
let mut acc = 0.0f64;
for i in 0..n {
let d = f64::from(p[i]) - f64::from(target[i]);
acc += d * d;
}
let mut result = Tensor::from_vec(vec![(acc * inv_n) as f32], &[1]);
if is_grad_enabled() && pred.requires_grad_enabled() {
result.requires_grad_(true);
let grad_fn = Arc::new(MseBackward {
pred: pred.clone(),
target: target.to_vec(),
});
result.set_grad_fn(grad_fn.clone());
with_graph(|graph| {
graph.register_tensor(pred.clone());
graph.record(result.id(), grad_fn, vec![pred.id()]);
});
}
contract_inv_mse_loss!(result.data());
Ok(result)
}