use serde::{Deserialize, Serialize};
use super::validation::{
deserialize_sample_rate, deserialize_submodels, deserialize_training, deserialize_weights,
};
#[derive(Deserialize, Serialize, Debug, Clone, PartialEq, Eq, Default)]
pub struct NamDate {
pub year: Option<i32>,
pub month: Option<i32>,
pub day: Option<i32>,
pub hour: Option<i32>,
pub minute: Option<i32>,
pub second: Option<i32>,
}
#[derive(Deserialize, Serialize, Debug, Clone, Default)]
pub struct NamMetadata {
pub date: Option<NamDate>,
pub name: Option<String>,
pub modeled_by: Option<String>,
pub gear_make: Option<String>,
pub gear_model: Option<String>,
pub gear_type: Option<String>,
pub tone_type: Option<String>,
#[serde(default, deserialize_with = "deserialize_training")]
pub training: Option<serde_json::Value>,
pub input_level_dbu: Option<f32>,
pub output_level_dbu: Option<f32>,
pub loudness: Option<f32>,
}
#[derive(Deserialize)]
#[serde(rename_all = "snake_case")]
struct NamLayerConfigHelper {
input_size: Option<usize>,
condition_size: Option<usize>,
head_size: Option<usize>,
channels: Option<usize>,
kernel_size: Option<usize>,
kernel_sizes: Option<Vec<usize>>,
dilations: Option<Vec<usize>>,
activation: Option<serde_json::Value>,
gated: Option<bool>,
head_bias: Option<bool>,
bottleneck: Option<usize>,
}
#[derive(Serialize, Debug, Clone, Default)]
pub struct NamLayerConfig {
pub input_size: Option<usize>,
pub condition_size: Option<usize>,
pub head_size: Option<usize>,
pub channels: Option<usize>,
pub kernel_size: Option<usize>,
pub kernel_sizes: Option<Vec<usize>>,
pub dilations: Option<Vec<usize>>,
pub activation: Option<String>,
pub gated: Option<bool>,
pub head_bias: Option<bool>,
pub bottleneck: Option<usize>,
#[serde(skip)]
pub layer_raw: Option<serde_json::Value>,
}
impl NamLayerConfig {
pub fn parse_activation_config(
&self,
num_layers: usize,
) -> Option<super::activation_parser::LayerActivationConfig> {
let raw = self.layer_raw.as_ref()?;
super::activation_parser::parse_layer_activations(raw, num_layers)
}
}
impl<'de> Deserialize<'de> for NamLayerConfig {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: serde::Deserializer<'de>,
{
let raw_value = serde_json::Value::deserialize(deserializer)?;
let helper: NamLayerConfigHelper =
serde_json::from_value(raw_value.clone()).map_err(serde::de::Error::custom)?;
let activation = match helper.activation {
Some(serde_json::Value::String(s)) => Some(s),
Some(serde_json::Value::Array(_)) => None,
Some(serde_json::Value::Object(_)) => None,
Some(serde_json::Value::Null) | None => None,
_ => {
return Err(serde::de::Error::custom(
"unsupported activation format: expected a string, array, object, or null",
));
}
};
Ok(NamLayerConfig {
input_size: helper.input_size,
condition_size: helper.condition_size,
head_size: helper.head_size,
channels: helper.channels,
kernel_size: helper.kernel_size,
kernel_sizes: helper.kernel_sizes,
dilations: helper.dilations,
activation,
gated: helper.gated,
head_bias: helper.head_bias,
bottleneck: helper.bottleneck,
layer_raw: Some(raw_value),
})
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum LinearImplementation {
#[default]
Auto,
Direct,
Fft,
}
impl LinearImplementation {}
impl std::str::FromStr for LinearImplementation {
type Err = ();
fn from_str(s: &str) -> Result<Self, Self::Err> {
let lower = s.to_lowercase();
match lower.as_str() {
"auto" => Ok(Self::Auto),
"direct" => Ok(Self::Direct),
"fft" | "partitioned_fft" | "partitioned-fft" => Ok(Self::Fft),
"legacy" | "old" => Ok(Self::Auto),
_ => Err(()),
}
}
}
#[cfg(test)]
#[path = "model_test.rs"]
mod tests;
#[derive(serde::Deserialize, serde::Serialize, Debug, Clone, Copy, PartialEq, Eq, Default)]
#[repr(u8)]
pub enum WeightsLayout {
#[default]
Original = 0,
GateMajorLstm = 1,
Interleaved4WaveNet = 2,
}
#[derive(Serialize, Debug, Clone, PartialEq, Default)]
pub struct HeadConfig {
pub channels: Option<usize>,
pub bias: Option<bool>,
pub out_channels: Option<usize>,
pub activation: Option<String>,
pub kernel_size: Option<usize>,
}
#[derive(Deserialize, Serialize, Debug, Clone, Default)]
pub struct NamConfig {
#[serde(default)]
pub layers: Vec<NamLayerConfig>,
#[serde(default)]
pub in_channels: Option<usize>,
#[serde(default)]
pub out_channels: Option<usize>,
pub head: Option<serde_json::Value>,
pub head_scale: Option<f32>,
pub num_layers: Option<usize>,
pub hidden_size: Option<usize>,
pub receptive_field: Option<usize>,
pub bias: Option<bool>,
#[serde(
default,
skip_serializing_if = "Option::is_none",
deserialize_with = "deserialize_submodels"
)]
pub submodels: Option<Vec<serde_json::Value>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub condition_dsp: Option<serde_json::Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub implementation: Option<String>,
#[serde(default, rename = "channels")]
pub conv_channels: Option<usize>,
#[serde(default, rename = "dilations")]
pub conv_dilations: Option<Vec<usize>>,
#[serde(default, rename = "batchnorm")]
pub conv_batchnorm: Option<bool>,
}
impl NamConfig {
pub fn parse_head(&self) -> Option<HeadConfig> {
let val = self.head.as_ref()?;
if val.is_null() || !val.is_object() {
return None;
}
let channels = val
.get("channels")
.and_then(|v| v.as_u64())
.and_then(|v| usize::try_from(v).ok());
let bias = val.get("bias").and_then(|v| v.as_bool());
let out_channels = val
.get("out_channels")
.and_then(|v| v.as_u64())
.and_then(|v| usize::try_from(v).ok());
let activation = val
.get("activation")
.and_then(|v| v.as_str())
.map(String::from);
let kernel_size = val
.get("kernel_size")
.and_then(|v| v.as_u64())
.and_then(|v| usize::try_from(v).ok());
Some(HeadConfig {
channels,
bias,
out_channels,
activation,
kernel_size,
})
}
}
#[derive(Deserialize, Serialize, Debug, Clone)]
pub struct NamModelData {
pub version: Option<String>,
pub architecture: String,
pub config: NamConfig,
#[serde(deserialize_with = "deserialize_weights")]
pub weights: Vec<f32>,
#[serde(default, deserialize_with = "deserialize_sample_rate")]
pub sample_rate: Option<f32>,
pub metadata: Option<NamMetadata>,
#[serde(skip)]
pub weights_layout: WeightsLayout,
}