use std::{
collections::HashMap,
fmt::{Display, Write},
sync::atomic::{AtomicUsize, Ordering},
};
use nove_tensor::{DType, Device, Tensor};
use crate::{Model, ModelError};
use super::{Conv2d, Conv2dBuilder, Linear, LinearBuilder, MaxPool2d, MaxPool2dBuilder, ReLU};
static ID: AtomicUsize = AtomicUsize::new(0);
#[derive(Debug, Clone)]
pub enum CNNLayer {
Conv2d(Conv2d),
ReLU(ReLU),
MaxPool2d(MaxPool2d),
Linear(Linear),
}
impl Model for CNNLayer {
type Input = Tensor;
type Output = Tensor;
fn forward(&mut self, input: Self::Input) -> Result<Self::Output, ModelError> {
match self {
CNNLayer::Conv2d(layer) => layer.forward(input),
CNNLayer::ReLU(layer) => layer.forward(input),
CNNLayer::MaxPool2d(layer) => layer.forward(input),
CNNLayer::Linear(layer) => layer.forward(input),
}
}
fn require_grad(&mut self, grad_enabled: bool) -> Result<(), ModelError> {
match self {
CNNLayer::Conv2d(layer) => layer.require_grad(grad_enabled),
CNNLayer::ReLU(layer) => layer.require_grad(grad_enabled),
CNNLayer::MaxPool2d(layer) => layer.require_grad(grad_enabled),
CNNLayer::Linear(layer) => layer.require_grad(grad_enabled),
}
}
fn to_device(&mut self, device: &Device) -> Result<(), ModelError> {
match self {
CNNLayer::Conv2d(layer) => layer.to_device(device),
CNNLayer::ReLU(layer) => layer.to_device(device),
CNNLayer::MaxPool2d(layer) => layer.to_device(device),
CNNLayer::Linear(layer) => layer.to_device(device),
}
}
fn to_dtype(&mut self, dtype: &DType) -> Result<(), ModelError> {
match self {
CNNLayer::Conv2d(layer) => layer.to_dtype(dtype),
CNNLayer::ReLU(layer) => layer.to_dtype(dtype),
CNNLayer::MaxPool2d(layer) => layer.to_dtype(dtype),
CNNLayer::Linear(layer) => layer.to_dtype(dtype),
}
}
fn parameters(&self) -> Result<Vec<Tensor>, ModelError> {
match self {
CNNLayer::Conv2d(layer) => layer.parameters(),
CNNLayer::ReLU(layer) => layer.parameters(),
CNNLayer::MaxPool2d(layer) => layer.parameters(),
CNNLayer::Linear(layer) => layer.parameters(),
}
}
fn named_parameters(&self) -> Result<HashMap<String, Tensor>, ModelError> {
match self {
CNNLayer::Conv2d(layer) => layer.named_parameters(),
CNNLayer::ReLU(layer) => layer.named_parameters(),
CNNLayer::MaxPool2d(layer) => layer.named_parameters(),
CNNLayer::Linear(layer) => layer.named_parameters(),
}
}
}
impl Display for CNNLayer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
CNNLayer::Conv2d(layer) => write!(f, "{}", layer),
CNNLayer::ReLU(layer) => write!(f, "{}", layer),
CNNLayer::MaxPool2d(layer) => write!(f, "{}", layer),
CNNLayer::Linear(layer) => write!(f, "{}", layer),
}
}
}
#[derive(Debug, Clone)]
pub struct CNN {
layers: Vec<CNNLayer>,
id: usize,
}
impl CNN {
pub fn layers(&self) -> &Vec<CNNLayer> {
&self.layers
}
pub fn layers_mut(&mut self) -> &mut Vec<CNNLayer> {
&mut self.layers
}
}
impl Model for CNN {
type Input = (Tensor, bool);
type Output = Tensor;
fn forward(&mut self, input: Self::Input) -> Result<Self::Output, ModelError> {
let (mut input, _training) = input;
for layer in &mut self.layers {
if let CNNLayer::Linear(_) = layer {
let shape = input.shape()?;
if shape.dims().len() == 4 {
let batch_size = shape.dims()[0];
let flattened_size = shape.dims()[1] * shape.dims()[2] * shape.dims()[3];
input =
input.reshape(&nove_tensor::Shape::from(&[batch_size, flattened_size]))?;
}
}
input = layer.forward(input)?;
}
Ok(input)
}
fn require_grad(&mut self, grad_enabled: bool) -> Result<(), ModelError> {
for layer in &mut self.layers {
layer.require_grad(grad_enabled)?;
}
Ok(())
}
fn to_device(&mut self, device: &Device) -> Result<(), ModelError> {
for layer in &mut self.layers {
layer.to_device(device)?;
}
Ok(())
}
fn to_dtype(&mut self, dtype: &DType) -> Result<(), ModelError> {
for layer in &mut self.layers {
layer.to_dtype(dtype)?;
}
Ok(())
}
fn parameters(&self) -> Result<Vec<Tensor>, ModelError> {
let mut params = Vec::new();
for layer in &self.layers {
params.extend(layer.parameters()?);
}
Ok(params)
}
fn named_parameters(&self) -> Result<HashMap<String, Tensor>, ModelError> {
let mut params = HashMap::new();
for (i, layer) in self.layers.iter().enumerate() {
let prefix = format!("cnn{}.", i);
for (name, tensor) in layer.named_parameters()? {
params.insert(format!("{}{}", prefix, name), tensor);
}
}
Ok(params)
}
}
impl Display for CNN {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let mut s = String::new();
writeln!(s, "cnn.{}(", self.id)?;
for layer in &self.layers {
writeln!(s, " {},", layer)?;
}
write!(s, ")")?;
write!(f, "{}", s)
}
}
#[derive(Debug, Clone)]
pub struct CNNConvBlock {
in_channels: usize,
out_channels: usize,
kernel_size: (usize, usize),
stride: usize,
padding: usize,
use_relu: bool,
use_pool: bool,
pool_kernel_size: (usize, usize),
pool_stride: (usize, usize),
}
impl CNNConvBlock {
pub fn new(in_channels: usize, out_channels: usize) -> Self {
Self {
in_channels,
out_channels,
kernel_size: (3, 3),
stride: 1,
padding: 1,
use_relu: true,
use_pool: false,
pool_kernel_size: (2, 2),
pool_stride: (2, 2),
}
}
pub fn kernel_size(mut self, kernel_size: (usize, usize)) -> Self {
self.kernel_size = kernel_size;
self
}
pub fn stride(mut self, stride: usize) -> Self {
self.stride = stride;
self
}
pub fn padding(mut self, padding: usize) -> Self {
self.padding = padding;
self
}
pub fn use_relu(mut self, use_relu: bool) -> Self {
self.use_relu = use_relu;
self
}
pub fn use_pool(mut self, use_pool: bool) -> Self {
self.use_pool = use_pool;
self
}
pub fn pool_kernel_size(mut self, pool_kernel_size: (usize, usize)) -> Self {
self.pool_kernel_size = pool_kernel_size;
self
}
pub fn pool_stride(mut self, pool_stride: (usize, usize)) -> Self {
self.pool_stride = pool_stride;
self
}
}
#[derive(Debug, Clone)]
pub struct CNNLinearBlock {
in_features: usize,
out_features: usize,
use_relu: bool,
bias_enabled: bool,
}
impl CNNLinearBlock {
pub fn new(in_features: usize, out_features: usize) -> Self {
Self {
in_features,
out_features,
use_relu: false,
bias_enabled: true,
}
}
pub fn use_relu(mut self, use_relu: bool) -> Self {
self.use_relu = use_relu;
self
}
pub fn bias_enabled(mut self, bias_enabled: bool) -> Self {
self.bias_enabled = bias_enabled;
self
}
}
pub struct CNNBuilder {
conv_blocks: Vec<CNNConvBlock>,
linear_blocks: Vec<CNNLinearBlock>,
device: Device,
dtype: DType,
grad_enabled: bool,
}
impl Default for CNNBuilder {
fn default() -> Self {
Self {
conv_blocks: Vec::new(),
linear_blocks: Vec::new(),
device: Device::cpu(),
dtype: DType::F32,
grad_enabled: true,
}
}
}
impl CNNBuilder {
pub fn conv_block(&mut self, block: CNNConvBlock) -> &mut Self {
self.conv_blocks.push(block);
self
}
pub fn linear_block(&mut self, block: CNNLinearBlock) -> &mut Self {
self.linear_blocks.push(block);
self
}
pub fn device(&mut self, device: Device) -> &mut Self {
self.device = device;
self
}
pub fn dtype(&mut self, dtype: DType) -> &mut Self {
self.dtype = dtype;
self
}
pub fn grad_enabled(&mut self, grad_enabled: bool) -> &mut Self {
self.grad_enabled = grad_enabled;
self
}
pub fn build(&self) -> Result<CNN, ModelError> {
let mut layers = Vec::new();
for block in &self.conv_blocks {
let conv = Conv2dBuilder::default()
.in_channels(block.in_channels)
.out_channels(block.out_channels)
.kernel_size(block.kernel_size)
.stride(block.stride)
.padding(block.padding)
.device(self.device.clone())
.dtype(self.dtype.clone())
.grad_enabled(self.grad_enabled)
.build()?;
layers.push(CNNLayer::Conv2d(conv));
if block.use_relu {
layers.push(CNNLayer::ReLU(ReLU::new()));
}
if block.use_pool {
let pool = MaxPool2dBuilder::default()
.kernel_size(block.pool_kernel_size)
.stride(block.pool_stride)
.build()?;
layers.push(CNNLayer::MaxPool2d(pool));
}
}
for block in &self.linear_blocks {
let linear = LinearBuilder::default()
.in_features(block.in_features)
.out_features(block.out_features)
.bias_enabled(block.bias_enabled)
.device(self.device.clone())
.dtype(self.dtype.clone())
.grad_enabled(self.grad_enabled)
.build()?;
layers.push(CNNLayer::Linear(linear));
if block.use_relu {
layers.push(CNNLayer::ReLU(ReLU::new()));
}
}
let id = ID.fetch_add(1, Ordering::Relaxed);
Ok(CNN { layers, id })
}
}