use crate::config::model::TensorConfig;
use anyhow::{Error as E, Result};
use serde_json::Value;
use std::collections::HashMap;
use tracing::debug;
#[derive(Debug, Clone, PartialEq)]
pub enum ComponentRole {
Embeddings, FfnPrefill, FfnInfer, FfnUnified, LmHead, Unknown, }
pub struct SchemaExtractor;
impl Default for SchemaExtractor {
fn default() -> Self {
Self::new()
}
}
impl SchemaExtractor {
pub fn new() -> Self {
Self
}
pub fn extract_inputs(&self, manifest: &Value) -> Result<HashMap<String, TensorConfig>> {
let mut inputs = HashMap::new();
if let Some(input_schema) = manifest
.get(0)
.and_then(|m| m.get("inputSchema").and_then(|s| s.as_array()))
{
debug!("📖 Parsing input schema with {} inputs", input_schema.len());
if let Some(function_inputs) = self.try_extract_function_inputs(manifest)? {
return Ok(function_inputs);
}
inputs = self.parse_tensor_configs(input_schema)?;
} else if manifest
.get("itemInfoEntries")
.and_then(|v| v.as_object())
.is_some()
{
debug!("📖 Parsing .mlpackage manifest (limited). Using minimal fallback");
} else {
debug!("📖 Unknown manifest format, no inputs extracted");
}
debug!("📖 Extracted {} input tensors", inputs.len());
Ok(inputs)
}
pub fn extract_outputs(&self, manifest: &Value) -> Result<HashMap<String, TensorConfig>> {
let mut outputs = HashMap::new();
if let Some(output_schema) = manifest
.get(0)
.and_then(|m| m.get("outputSchema").and_then(|s| s.as_array()))
{
debug!(
"📖 Parsing output schema with {} outputs",
output_schema.len()
);
outputs = self.parse_tensor_configs(output_schema)?;
if self.has_empty_shapes(&outputs) {
self.backfill_from_function_schemas(manifest, &mut outputs)?;
}
} else if let Some(functions) = manifest
.get(0)
.and_then(|m| m.get("functions").and_then(|f| f.as_array()))
{
debug!(
"📖 Parsing outputs from functions schema ({} functions)",
functions.len()
);
outputs = self.extract_outputs_from_functions(functions)?;
} else {
debug!("📖 Unknown manifest format, no outputs extracted");
}
debug!("📖 Extracted {} output tensors", outputs.len());
Ok(outputs)
}
pub fn parse_tensor_configs(&self, schema: &[Value]) -> Result<HashMap<String, TensorConfig>> {
let mut configs = HashMap::new();
for tensor_def in schema {
if let Some(tensor_config) = self.parse_tensor_definition(tensor_def)? {
configs.insert(tensor_config.name.clone(), tensor_config);
}
}
debug!("📖 Extracted {} tensor configs", configs.len());
Ok(configs)
}
fn try_extract_function_inputs(
&self,
manifest: &Value,
) -> Result<Option<HashMap<String, TensorConfig>>> {
let functions = manifest
.get(0)
.and_then(|m| m.get("functions").and_then(|f| f.as_array()));
if let Some(funcs) = functions {
for prefer in ["prefill", "infer"] {
for function in funcs {
if let Some(func_name) = function.get("name").and_then(|n| n.as_str()) {
if func_name == prefer {
if let Some(input_schema) =
function.get("inputSchema").and_then(|s| s.as_array())
{
debug!(
"📖 Using {} function input schema with {} inputs",
prefer,
input_schema.len()
);
return Ok(Some(self.parse_tensor_configs(input_schema)?));
}
}
}
}
}
}
Ok(None)
}
fn has_empty_shapes(&self, outputs: &HashMap<String, TensorConfig>) -> bool {
outputs.values().any(|tensor| tensor.shape.is_empty())
}
fn backfill_from_function_schemas(
&self,
manifest: &Value,
outputs: &mut HashMap<String, TensorConfig>,
) -> Result<()> {
let functions = manifest
.get(0)
.and_then(|m| m.get("functions").and_then(|f| f.as_array()));
if let Some(funcs) = functions {
for prefer in ["prefill", "infer"] {
for function in funcs {
if let Some(func_name) = function.get("name").and_then(|n| n.as_str()) {
if func_name == prefer {
if let Some(func_output_schema) =
function.get("outputSchema").and_then(|s| s.as_array())
{
self.merge_function_outputs(func_output_schema, outputs)?;
}
}
}
}
}
}
Ok(())
}
fn merge_function_outputs(
&self,
func_output_schema: &[Value],
outputs: &mut HashMap<String, TensorConfig>,
) -> Result<()> {
let function_outputs = self.parse_tensor_configs(func_output_schema)?;
for (name, tensor_config) in function_outputs {
let entry = outputs.entry(name.clone()).or_insert(TensorConfig {
name: name.clone(),
shape: vec![],
data_type: tensor_config.data_type.clone(),
});
if entry.shape.is_empty() {
entry.shape = tensor_config.shape;
entry.data_type = tensor_config.data_type;
}
}
Ok(())
}
fn extract_outputs_from_functions(
&self,
functions: &[Value],
) -> Result<HashMap<String, TensorConfig>> {
let mut outputs = HashMap::new();
for function in functions {
if let Some(func_output_schema) =
function.get("outputSchema").and_then(|s| s.as_array())
{
let function_outputs = self.parse_tensor_configs(func_output_schema)?;
outputs.extend(function_outputs);
}
}
Ok(outputs)
}
fn parse_tensor_definition(&self, tensor_def: &Value) -> Result<Option<TensorConfig>> {
let (name, shape_str, data_type) = match (
tensor_def.get("name").and_then(|n| n.as_str()),
tensor_def.get("shape").and_then(|s| s.as_str()),
tensor_def.get("dataType").and_then(|d| d.as_str()),
) {
(Some(n), Some(s), Some(d)) => (n, s, d),
_ => return Ok(None),
};
let shape = self.parse_shape_string(shape_str)?;
debug!(" Tensor: {} -> {:?} ({})", name, shape, data_type);
Ok(Some(TensorConfig {
name: name.to_string(),
shape,
data_type: data_type.to_uppercase(),
}))
}
#[allow(dead_code)]
fn extract_tensor_shape(&self, input: &Value) -> Result<Vec<usize>> {
if let Some(enum_val) = input.get("enumeratedShapes") {
if let Some(shapes) = self.parse_enumerated_shapes(enum_val)? {
return Ok(shapes);
}
}
if let Some(shape_str) = input.get("shape").and_then(|s| s.as_str()) {
return self.parse_shape_string(shape_str);
}
Ok(vec![])
}
#[allow(dead_code)]
fn parse_enumerated_shapes(&self, enum_val: &Value) -> Result<Option<Vec<usize>>> {
if let Some(enum_str) = enum_val.as_str() {
match serde_json::from_str::<Vec<Vec<usize>>>(enum_str) {
Ok(mut shapes) => {
shapes.sort_by(|a, b| {
let a_size: usize = a.iter().product();
let b_size: usize = b.iter().product();
a_size.cmp(&b_size)
});
return Ok(shapes.last().cloned());
}
Err(err) => {
debug!("⚠️ Failed to parse enumeratedShapes: {}", err);
}
}
} else if let Some(enum_arr) = enum_val.as_array() {
let candidates = self.parse_enumerated_array(enum_arr)?;
if !candidates.is_empty() {
let mut sorted_candidates = candidates;
sorted_candidates.sort_by(|a, b| {
let a_size: usize = a.iter().product();
let b_size: usize = b.iter().product();
a_size.cmp(&b_size)
});
return Ok(sorted_candidates.last().cloned());
}
}
Ok(None)
}
#[allow(dead_code)]
fn parse_enumerated_array(&self, enum_arr: &[Value]) -> Result<Vec<Vec<usize>>> {
let mut candidates = Vec::new();
for item in enum_arr {
if let Some(s) = item.as_str() {
if let Ok(v) = serde_json::from_str::<Vec<usize>>(s) {
candidates.push(v);
}
} else if let Some(arr) = item.as_array() {
let mut v = Vec::new();
for d in arr {
if let Some(u) = d.as_u64() {
v.push(u as usize);
}
}
if !v.is_empty() {
candidates.push(v);
}
}
}
Ok(candidates)
}
fn parse_shape_string(&self, shape_str: &str) -> Result<Vec<usize>> {
let trimmed = shape_str.trim_start_matches('[').trim_end_matches(']');
if trimmed.is_empty() {
return Ok(vec![]);
}
let mut dims = Vec::new();
for dim_str in trimmed.split(',') {
match dim_str.trim().parse::<usize>() {
Ok(dim) => dims.push(dim),
Err(_) => return Err(E::msg("Failed to parse tensor dimension")),
}
}
Ok(dims)
}
pub fn detect_component_role(
&self,
inputs: &HashMap<String, TensorConfig>,
outputs: &HashMap<String, TensorConfig>,
) -> ComponentRole {
debug!("🔍 Analyzing tensor signatures for component role detection");
debug!(" Inputs: {:?}", inputs.keys().collect::<Vec<_>>());
debug!(" Outputs: {:?}", outputs.keys().collect::<Vec<_>>());
if inputs.is_empty() && outputs.is_empty() {
debug!("⚠️ No tensor information available for component role detection");
return ComponentRole::Unknown;
}
if inputs.contains_key("input_ids") && outputs.contains_key("hidden_states") {
debug!("✅ Detected EMBEDDINGS component (input_ids -> hidden_states)");
return ComponentRole::Embeddings;
}
if outputs.keys().any(|k| k.starts_with("logits")) {
let logit_count = outputs.keys().filter(|k| k.starts_with("logits")).count();
debug!(
"✅ Detected LM_HEAD component ({} logits outputs)",
logit_count
);
return ComponentRole::LmHead;
}
if inputs.contains_key("hidden_states") && outputs.contains_key("output_hidden_states") {
let has_causal_mask = inputs.contains_key("causal_mask");
let has_update_mask = inputs.contains_key("update_mask");
if has_update_mask && has_causal_mask {
debug!("✅ Detected FFN_INFER component (has update_mask + causal_mask)");
return ComponentRole::FfnInfer;
} else if has_causal_mask && !has_update_mask {
debug!("✅ Detected FFN component with causal_mask (prefill/infer ambiguous - needs filename fallback)");
return ComponentRole::Unknown;
} else {
debug!("✅ Detected FFN_UNIFIED component (no distinctive masks)");
return ComponentRole::FfnUnified;
}
}
debug!("❓ Could not determine component role from tensor signatures");
ComponentRole::Unknown
}
pub fn detect_component_role_from_filename(&self, filename: &str) -> ComponentRole {
let filename_lower = filename.to_lowercase();
if filename_lower.contains("embedding") {
debug!("🏷️ Filename-based detection: EMBEDDINGS ({})", filename);
return ComponentRole::Embeddings;
}
if filename_lower.contains("lm_head") || filename_lower.contains("lmhead") {
debug!("🏷️ Filename-based detection: LM_HEAD ({})", filename);
return ComponentRole::LmHead;
}
if filename_lower.contains("ffn_pf") {
debug!(
"🏷️ Filename-based detection: FFN_UNIFIED (ANEMLL FFN_PF pattern: {})",
filename
);
return ComponentRole::FfnUnified;
}
if filename_lower.contains("prefill") && !filename_lower.contains("ffn_pf") {
debug!(
"🏷️ Filename-based detection: FFN_PREFILL (split architecture: {})",
filename
);
return ComponentRole::FfnPrefill;
}
if filename_lower.contains("infer")
|| (filename_lower.contains("ffn_chunk")
&& !filename_lower.contains("ffn_pf")
&& !filename_lower.contains("prefill"))
{
debug!(
"🏷️ Filename-based detection: FFN_INFER (split architecture: {})",
filename
);
return ComponentRole::FfnInfer;
}
debug!("🏷️ Filename-based detection: UNKNOWN ({})", filename);
ComponentRole::Unknown
}
pub fn calculate_vocab_size_from_logits(
&self,
outputs: &HashMap<String, TensorConfig>,
) -> Option<usize> {
let mut total_vocab_size = 0;
let mut logits_found = false;
if let Some(logits_tensor) = outputs.get("logits") {
if let Some(&vocab_dim) = logits_tensor.shape.last() {
debug!("📊 Single logits tensor found, vocab_size: {}", vocab_dim);
return Some(vocab_dim);
}
}
for (name, tensor) in outputs {
if name.starts_with("logits") && name != "logits" {
if let Some(&chunk_size) = tensor.shape.last() {
total_vocab_size += chunk_size;
logits_found = true;
debug!("📊 Found logits chunk {}: size {}", name, chunk_size);
}
}
}
if logits_found {
debug!(
"📊 Total vocab size from {} chunks: {}",
outputs
.keys()
.filter(|k| k.starts_with("logits") && *k != "logits")
.count(),
total_vocab_size
);
Some(total_vocab_size)
} else {
None
}
}
}