use std::fmt::Display;
use crate::Parameter;
use nove_tensor::{Device, TensorError};
use thiserror::Error;
pub mod safetensors;
#[derive(Error, Debug)]
pub enum ParamStoreError {
#[error("IO error: {0}")]
IoError(#[from] std::io::Error),
#[error("Tensor error: {0}")]
TensorError(#[from] TensorError),
#[error("RwLock poisoned: {0}")]
RwLockPoisoned(String),
#[error("Other error: {0}")]
OtherError(String),
}
impl<T> From<std::sync::PoisonError<T>> for ParamStoreError {
fn from(error: std::sync::PoisonError<T>) -> Self {
ParamStoreError::RwLockPoisoned(error.to_string())
}
}
pub trait ParamStore: Display + Clone {
fn new(name: &str) -> Result<Self, ParamStoreError>;
fn set_name(&self, name: &str) -> Result<(), ParamStoreError>;
fn name(&self) -> Result<String, ParamStoreError>;
fn save(&self, folder_path: &str) -> Result<(), ParamStoreError>;
fn load<F>(
&self,
folder_path: &str,
device: &Device,
process_fn: F,
) -> Result<(), ParamStoreError>
where
F: FnMut(&str, &Self) -> Result<(), ParamStoreError>;
fn set_module(&self, module: Self) -> Result<(), ParamStoreError>;
fn modules(&self) -> Result<Vec<Self>, ParamStoreError>;
fn set_parameter(&self, param: Parameter) -> Result<(), ParamStoreError>;
fn parameters(&self) -> Result<Vec<Parameter>, ParamStoreError>;
}