use serde_json::Value;
use crate::types::DiffResult;
pub fn analyze_memory_usage_changes(
old_model: &Value,
new_model: &Value,
results: &mut Vec<DiffResult>,
) {
let old_memory = calculate_model_memory_usage(old_model);
let new_memory = calculate_model_memory_usage(new_model);
if old_memory != new_memory {
let memory_change = new_memory as f64 - old_memory as f64;
let memory_change_percent = if old_memory > 0 {
(memory_change / old_memory as f64) * 100.0
} else {
0.0
};
let _memory_analysis =
format!("memory: {old_memory} → {new_memory} bytes ({memory_change_percent:+.1}%)");
results.push(DiffResult::ModelArchitectureChanged(
"memory_analysis".to_string(),
format!("memory_usage: {old_memory} bytes"),
format!("memory_usage: {new_memory} bytes"),
));
if memory_change.abs() > 1024.0 {
let breakdown = create_memory_breakdown(old_model, new_model);
if !breakdown.is_empty() {
results.push(DiffResult::ModelArchitectureChanged(
"memory_breakdown".to_string(),
"previous".to_string(),
breakdown,
));
}
}
}
}
pub(crate) fn calculate_model_memory_usage(model: &Value) -> usize {
match model {
Value::Object(obj) => {
let mut total_memory = 0;
total_memory += std::mem::size_of::<serde_json::Map<String, Value>>();
for (key, value) in obj {
total_memory += key.len();
total_memory += calculate_value_memory(value);
}
total_memory
}
_ => calculate_value_memory(model),
}
}
pub(crate) fn calculate_value_memory(value: &Value) -> usize {
match value {
Value::Null => std::mem::size_of::<Value>(),
Value::Bool(_) => std::mem::size_of::<bool>(),
Value::Number(_) => std::mem::size_of::<f64>(), Value::String(s) => s.len() + std::mem::size_of::<String>(),
Value::Array(arr) => {
let mut size = std::mem::size_of::<Vec<Value>>();
for elem in arr {
size += calculate_value_memory(elem);
}
size
}
Value::Object(obj) => {
let mut size = std::mem::size_of::<serde_json::Map<String, Value>>();
for (key, val) in obj {
size += key.len() + calculate_value_memory(val);
}
size
}
}
}
pub(crate) fn create_memory_breakdown(old_model: &Value, new_model: &Value) -> String {
let mut breakdown = Vec::new();
if let (Value::Object(old_obj), Value::Object(new_obj)) = (old_model, new_model) {
let old_tensor_memory = calculate_tensor_memory(old_obj);
let new_tensor_memory = calculate_tensor_memory(new_obj);
if old_tensor_memory != new_tensor_memory {
let change = new_tensor_memory as i64 - old_tensor_memory as i64;
breakdown.push(format!(
"tensors: {change:+} bytes ({old_tensor_memory} → {new_tensor_memory})"
));
}
let old_meta_memory = calculate_metadata_memory(old_obj);
let new_meta_memory = calculate_metadata_memory(new_obj);
if old_meta_memory != new_meta_memory {
let change = new_meta_memory as i64 - old_meta_memory as i64;
breakdown.push(format!(
"metadata: {change:+} bytes ({old_meta_memory} → {new_meta_memory})"
));
}
}
breakdown.join(", ")
}
pub(crate) fn calculate_tensor_memory(obj: &serde_json::Map<String, Value>) -> usize {
let mut tensor_memory = 0;
for (key, value) in obj {
if key.contains("weight") || key.contains("bias") || key.contains("data") {
if let Value::Object(tensor_obj) = value {
if let Some(shape_value) = tensor_obj.get("shape") {
if let Value::Array(shape_arr) = shape_value {
let element_count: usize = shape_arr
.iter()
.filter_map(|v| v.as_u64())
.map(|x| x as usize)
.product();
let dtype_size = if let Some(dtype) = tensor_obj.get("dtype") {
estimate_dtype_size(dtype)
} else {
4
};
tensor_memory += element_count * dtype_size;
}
}
} else {
tensor_memory += calculate_value_memory(value);
}
}
}
tensor_memory
}
pub(crate) fn calculate_metadata_memory(obj: &serde_json::Map<String, Value>) -> usize {
let mut meta_memory = 0;
for (key, value) in obj {
if !key.contains("weight") && !key.contains("bias") && !key.contains("data") {
meta_memory += key.len() + calculate_value_memory(value);
}
}
meta_memory
}
pub(crate) fn estimate_dtype_size(dtype: &Value) -> usize {
if let Value::String(dtype_str) = dtype {
match dtype_str.to_lowercase().as_str() {
s if s.contains("float64") || s.contains("f64") => 8,
s if s.contains("float32") || s.contains("f32") => 4,
s if s.contains("float16") || s.contains("f16") => 2,
s if s.contains("int64") || s.contains("i64") => 8,
s if s.contains("int32") || s.contains("i32") => 4,
s if s.contains("int16") || s.contains("i16") => 2,
s if s.contains("int8") || s.contains("i8") => 1,
s if s.contains("uint64") || s.contains("u64") => 8,
s if s.contains("uint32") || s.contains("u32") => 4,
s if s.contains("uint16") || s.contains("u16") => 2,
s if s.contains("uint8") || s.contains("u8") => 1,
s if s.contains("bool") => 1,
_ => 4, }
} else {
4 }
}