#[provable_contracts_macros::contract(
"setfit-encoder-conformance-v1",
equation = "l2_normalize_rows"
)]
pub fn l2_normalize_rows(x: &Tensor, eps: f32) -> Result<Tensor, OpError> {
let shape = x.shape();
if shape.len() != 2 {
return Err(OpError::ShapeMismatch {
expected: vec![0, 0],
got: shape.to_vec(),
});
}
let batch = shape[0];
let hidden = 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 xd = x.data();
if let Some(position) = xd.iter().position(|v| !v.is_finite()) {
return Err(OpError::NonFiniteInput { position });
}
contract_pre_l2_normalize_rows!(xd);
let mut out = vec![0.0f32; batch * hidden];
let mut norms = Vec::with_capacity(batch);
for row in 0..batch {
let base = row * hidden;
let sumsq: f64 = xd[base..base + hidden]
.iter()
.map(|&v| f64::from(v) * f64::from(v))
.sum();
let n = sumsq.sqrt() as f32;
let d = if n > eps { n } else { eps };
let inv = 1.0 / f64::from(d);
for j in 0..hidden {
out[base + j] = (f64::from(xd[base + j]) * inv) as f32;
}
norms.push(n);
}
let mut result = Tensor::from_vec(out, &[batch, hidden]);
if is_grad_enabled() && x.requires_grad_enabled() {
result.requires_grad_(true);
let grad_fn = Arc::new(L2NormalizeRowsBackward {
output: result.clone(),
norms,
eps,
batch,
hidden,
});
result.set_grad_fn(grad_fn.clone());
with_graph(|graph| {
graph.register_tensor(x.clone());
graph.record(result.id(), grad_fn, vec![x.id()]);
});
}
contract_post_l2_normalize_rows!(result.data());
Ok(result)
}