MCPcopy Create free account
hub / github.com/evilsocket/cake / read_layer_tensor_details

Function read_layer_tensor_details

cake-core/src/cake/sharding/api/ui.rs:207–302  ·  view source on GitHub ↗

Read per-layer tensor metadata from safetensors file headers (via mmap). Only the header bytes are paged in — tensor data is never accessed.

(
    data_path: &std::path::Path,
    num_layers: usize,
    layer_prefix: &str,
)

Source from the content-addressed store, hash-verified

205/// Read per-layer tensor metadata from safetensors file headers (via mmap).
206/// Only the header bytes are paged in — tensor data is never accessed.
207fn read_layer_tensor_details(
208 data_path: &std::path::Path,
209 num_layers: usize,
210 layer_prefix: &str,
211) -> HashMap<String, LayerDetail> {
212 // Collect safetensors shard files
213 let index_path = data_path.join("model.safetensors.index.json");
214 let shard_files: Vec<std::path::PathBuf> = if let Ok(data) =
215 std::fs::read_to_string(&index_path)
216 {
217 if let Ok(json) = serde_json::from_str::<serde_json::Value>(&data) {
218 if let Some(weight_map) = json.get("weight_map").and_then(|v| v.as_object()) {
219 let shards: HashSet<&str> =
220 weight_map.values().filter_map(|v| v.as_str()).collect();
221 shards.iter().map(|s| data_path.join(s)).collect()
222 } else {
223 vec![]
224 }
225 } else {
226 vec![]
227 }
228 } else {
229 let single = data_path.join("model.safetensors");
230 if single.exists() {
231 vec![single]
232 } else {
233 vec![]
234 }
235 };
236
237 // Read headers from each shard via mmap (parallel)
238 let all_tensors: Vec<(String, u64, String, Vec<usize>)> = shard_files
239 .par_iter()
240 .flat_map(|shard_path| {
241 let mut entries = Vec::new();
242 let file = match std::fs::File::open(shard_path) {
243 Ok(f) => f,
244 Err(_) => return entries,
245 };
246 let buffer = match unsafe { memmap2::MmapOptions::new().map(&file) } {
247 Ok(b) => b,
248 Err(_) => return entries,
249 };
250 let tensors = match SafeTensors::deserialize(&buffer) {
251 Ok(t) => t,
252 Err(_) => return entries,
253 };
254
255 for name in tensors.names() {
256 if let Ok(tv) = tensors.tensor(name) {
257 let shape: Vec<usize> = tv.shape().to_vec();
258 let size_bytes = tv.data().len() as u64;
259 let dtype = format!("{:?}", tv.dtype());
260 entries.push((name.to_string(), size_bytes, dtype, shape));
261 }
262 }
263 entries
264 })

Callers 1

topologyFunction · 0.85

Calls 6

to_vecMethod · 0.80
shapeMethod · 0.80
dataMethod · 0.80
getMethod · 0.45
pushMethod · 0.45
cloneMethod · 0.45

Tested by

no test coverage detected