use crate::models::a2::activations::ActivationType;
use crate::models::a2::gating::GatingMode;
#[derive(Debug, Clone, PartialEq)]
pub struct LayerActivationConfig {
pub activations: Vec<ActivationType>,
pub gating_modes: Vec<GatingMode>,
pub secondary_activations: Vec<Option<ActivationType>>,
}
pub fn parse_layer_activations(
raw: &serde_json::Value,
num_layers: usize,
) -> Option<LayerActivationConfig> {
let activations = parse_activations_from_json(raw, num_layers)?;
let gating_modes = parse_gating_modes_from_json(raw, num_layers);
let secondary_activations = parse_secondary_activations_from_json(raw, num_layers);
Some(LayerActivationConfig {
activations,
gating_modes,
secondary_activations,
})
}
pub fn parse_activations_from_json(
raw: &serde_json::Value,
num_layers: usize,
) -> Option<Vec<ActivationType>> {
if let Some(obj) = raw.get("activation").and_then(|v| v.as_object()) {
let at: ActivationType =
serde_json::from_value(serde_json::Value::Object(obj.clone())).ok()?;
return Some(vec![at; num_layers]);
}
let arr = raw.get("activation").and_then(|v| v.as_array())?;
if arr.len() != num_layers {
return None;
}
let mut out = Vec::with_capacity(num_layers);
for entry in arr {
let at: ActivationType = serde_json::from_value(entry.clone()).ok()?;
out.push(at);
}
Some(out)
}
pub fn parse_gating_modes_from_json(raw: &serde_json::Value, num_layers: usize) -> Vec<GatingMode> {
if let Some(s) = raw.get("gating_mode").and_then(|v| v.as_str()) {
let mode = match s {
"gated" => GatingMode::Gated,
"blended" => GatingMode::Blended,
_ => GatingMode::None,
};
return vec![mode; num_layers];
}
let arr = match raw.get("gating_mode") {
None | Some(serde_json::Value::Null) => {
return vec![GatingMode::None; num_layers];
}
Some(v) => match v.as_array() {
Some(a) if a.len() == num_layers => a,
_ => return vec![GatingMode::None; num_layers],
},
};
let mut out = Vec::with_capacity(num_layers);
for entry in arr {
let mode = match entry.as_str() {
Some("gated") => GatingMode::Gated,
Some("blended") => GatingMode::Blended,
_ => GatingMode::None,
};
out.push(mode);
}
out
}
pub fn parse_secondary_activations_from_json(
raw: &serde_json::Value,
num_layers: usize,
) -> Vec<Option<ActivationType>> {
if let Some(v) = raw.get("secondary_activation") {
if let Some(s) = v.as_str() {
if s.is_empty() || s.eq_ignore_ascii_case("none") {
return vec![None; num_layers];
}
let at = serde_json::from_value(serde_json::Value::String(s.to_string())).ok();
return vec![at; num_layers];
}
if v.is_object() {
let at = serde_json::from_value(v.clone()).ok();
return vec![at; num_layers];
}
}
let arr = match raw.get("secondary_activation") {
None | Some(serde_json::Value::Null) => {
return vec![None; num_layers];
}
Some(v) => match v.as_array() {
Some(a) if a.len() == num_layers => a,
_ => return vec![None; num_layers],
},
};
let mut out = Vec::with_capacity(num_layers);
for entry in arr {
if entry.is_null() {
out.push(None);
} else {
let at = serde_json::from_value(entry.clone()).ok();
out.push(at);
}
}
out
}
#[cfg(test)]
#[path = "activation_parser_test.rs"]
mod tests;