use teeny_core::{
dtype::Float,
graph::SymTensor,
name_scope::name_scope,
nn::{Layer, activation::sigmoid::Silu, batchnorm::BatchNorm2d, conv2d::Conv2d},
};
pub fn conv<D: Float>(
c_in: usize,
c_out: usize,
k: usize,
s: usize,
) -> impl Fn(SymTensor) -> SymTensor {
let p = k / 2;
let conv2d = Conv2d::<D, SymTensor, SymTensor, 4>::new(c_in, c_out, (k, k), (s, s), (p, p), false);
let bn = BatchNorm2d::<D, SymTensor, SymTensor, 4>::new(c_out).with_eps(0.001);
let act = Silu::<D, SymTensor, 4>::new();
move |x: SymTensor| {
let x = { let _g = name_scope("conv"); conv2d.call(x) };
let x = { let _g = name_scope("bn"); bn.call(x) };
act.call(x)
}
}
pub fn conv_bn<D: Float>(
c_in: usize,
c_out: usize,
k: usize,
s: usize,
g: usize,
) -> impl Fn(SymTensor) -> SymTensor {
let p = k / 2;
let conv2d = Conv2d::<D, SymTensor, SymTensor, 4>::new_grouped(
c_in, c_out, (k, k), (s, s), (p, p), false, g,
);
let bn = BatchNorm2d::<D, SymTensor, SymTensor, 4>::new(c_out).with_eps(0.001);
move |x: SymTensor| {
let x = { let _g = name_scope("conv"); conv2d.call(x) };
{ let _g = name_scope("bn"); bn.call(x) }
}
}
pub fn dwconv<D: Float>(c: usize, k: usize, s: usize) -> impl Fn(SymTensor) -> SymTensor {
let p = k / 2;
let conv2d = Conv2d::<D, SymTensor, SymTensor, 4>::new_grouped(c, c, (k, k), (s, s), (p, p), false, c);
let bn = BatchNorm2d::<D, SymTensor, SymTensor, 4>::new(c).with_eps(0.001);
let act = Silu::<D, SymTensor, 4>::new();
move |x: SymTensor| {
let x = { let _g = name_scope("conv"); conv2d.call(x) };
let x = { let _g = name_scope("bn"); bn.call(x) };
act.call(x)
}
}
pub fn conv_plain<D: Float>(
c_in: usize,
c_out: usize,
k: usize,
s: usize,
) -> impl Fn(SymTensor) -> SymTensor {
let p = k / 2;
let layer = Conv2d::<D, SymTensor, SymTensor, 4>::new(c_in, c_out, (k, k), (s, s), (p, p), true);
move |x| layer.call(x)
}