use kopitiam_core::Result;
use kopitiam_tensor::Tensor;
pub(crate) fn linear(x: &Tensor, weight: &Tensor, bias: Option<&Tensor>) -> Result<Tensor> {
let y = if kopitiam_tensor::has_fused_matmul_kernel(weight.dtype()) {
x.quantized_matmul(weight)?
} else {
let weight_t = weight.transpose(0, 1)?;
x.matmul(&weight_t)?
};
match bias {
Some(b) => y.add(b),
None => Ok(y),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn linear_without_bias_matches_matmul_by_the_transposed_weight() {
let x = Tensor::from_f32(vec![1.0, 2.0], [1, 2]).unwrap();
let w = Tensor::from_f32(vec![1.0, 0.0, 0.0, 1.0], [2, 2]).unwrap();
let y = linear(&x, &w, None).unwrap();
assert_eq!(y.to_vec_f32().unwrap(), vec![1.0, 2.0]);
}
#[test]
fn linear_applies_bias_after_the_matmul() {
let x = Tensor::from_f32(vec![1.0, 2.0], [1, 2]).unwrap();
let w = Tensor::from_f32(vec![1.0, 0.0, 0.0, 1.0], [2, 2]).unwrap();
let b = Tensor::from_f32(vec![10.0, 20.0], [2]).unwrap();
let y = linear(&x, &w, Some(&b)).unwrap();
assert_eq!(y.to_vec_f32().unwrap(), vec![11.0, 22.0]);
}
#[test]
fn linear_matches_hand_computation_for_a_non_trivial_weight() {
let x = Tensor::from_f32(vec![1.0, 2.0], [1, 2]).unwrap();
let w = Tensor::from_f32(vec![2.0, 0.0, 1.0, 3.0], [2, 2]).unwrap();
let y = linear(&x, &w, None).unwrap();
assert_eq!(y.to_vec_f32().unwrap(), vec![2.0, 7.0]);
}
#[test]
fn linear_dispatches_to_quantized_matmul_for_a_quantized_weight() {
let in_features = 32;
let row_values = [1.0f32, 2.0];
let mut bytes = Vec::new();
for &row_value in &row_values {
let d = row_value / 127.0;
bytes.extend_from_slice(&kopitiam_tensor::f32_to_f16(d).to_le_bytes());
bytes.extend(std::iter::repeat_n(127u8, in_features));
}
let weight = Tensor::from_quantized(kopitiam_tensor::DType::Q8_0, bytes, [row_values.len(), in_features]).unwrap();
let x = Tensor::from_f32(vec![1.0; in_features], [1, in_features]).unwrap();
let y = linear(&x, &weight, None).unwrap();
let out = y.to_vec_f32().unwrap();
assert_eq!(out.len(), row_values.len());
for (got, &row_value) in out.iter().zip(&row_values) {
let expected = in_features as f32 * row_value;
assert!((got - expected).abs() < 1e-2, "expected {expected}, got {got}");
}
}
#[test]
fn linear_applies_bias_after_a_quantized_matmul_too() {
let in_features = 32;
let d = 1.0f32 / 127.0;
let mut bytes = kopitiam_tensor::f32_to_f16(d).to_le_bytes().to_vec();
bytes.extend(std::iter::repeat_n(127u8, in_features));
let weight = Tensor::from_quantized(kopitiam_tensor::DType::Q8_0, bytes, [1, in_features]).unwrap();
let x = Tensor::from_f32(vec![1.0; in_features], [1, in_features]).unwrap();
let bias = Tensor::from_f32(vec![100.0], [1]).unwrap();
let y = linear(&x, &weight, Some(&bias)).unwrap();
let out = y.to_vec_f32().unwrap();
assert_eq!(out.len(), 1);
assert!((out[0] - (in_features as f32 + 100.0)).abs() < 1e-2, "got {}", out[0]);
}
}