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

Method new_stacked

cake-core/src/models/common/disk_expert_provider.rs:192–330  ·  view source on GitHub ↗

Create a disk-backed expert provider for the stacked switch_mlp format. In this format, expert weights are stored as 3D stacked tensors: `"{prefix}.gate_proj.weight"` with shape `(num_experts, intermediate, hidden)`. Individual expert slices are read via `tensor_slice_bytes()`.

(
        storage: Arc<dyn TensorStorageProvider>,
        layer_prefix: String,
        num_experts: usize,
        device: Device,
        dtype: DType,
    )

Source from the content-addressed store, hash-verified

190 /// `"{prefix}.gate_proj.weight"` with shape `(num_experts, intermediate, hidden)`.
191 /// Individual expert slices are read via `tensor_slice_bytes()`.
192 pub fn new_stacked(
193 storage: Arc<dyn TensorStorageProvider>,
194 layer_prefix: String,
195 num_experts: usize,
196 device: Device,
197 dtype: DType,
198 ) -> Self {
199 let expert_names: Vec<ExpertNames> = (0..num_experts)
200 .map(|_| {
201 ExpertNames {
202 gate_proj: format!("{layer_prefix}.gate_proj.weight"),
203 up_proj: format!("{layer_prefix}.up_proj.weight"),
204 down_proj: format!("{layer_prefix}.down_proj.weight"),
205 }
206 })
207 .collect();
208 let needs_device_transfer = !device.is_cpu();
209 // Pre-compute stacked metadata from the 3D tensor shapes
210 let stacked_meta = {
211 let gate_meta = storage.tensor_meta(&expert_names[0].gate_proj);
212 let up_meta = storage.tensor_meta(&expert_names[0].up_proj);
213 let down_meta = storage.tensor_meta(&expert_names[0].down_proj);
214 // Check for affine 4-bit quantization (scales + biases companions)
215 let scales_name = expert_names[0].gate_proj.replace(".weight", ".scales");
216 let is_affine = storage.has_tensor(&scales_name);
217 match (gate_meta, up_meta, down_meta) {
218 (Some((gdt, gs)), Some((_, us)), Some((_, ds))) if gs.len() == 3 && us.len() == 3 && ds.len() == 3 => {
219 let dtype_size = gdt.size_in_bytes();
220 let make_proj = |wshape: &[usize], proj_name: &str| -> StackedProjMeta {
221 let w_stride = wshape[1] * wshape[2] * dtype_size;
222 let (s_stride, s_shape) = if is_affine {
223 let sn = proj_name.replace(".weight", ".scales");
224 if let Some((sdt, ss)) = storage.tensor_meta(&sn) {
225 let ss_size = sdt.size_in_bytes();
226 (ss[1] * ss[2] * ss_size, [ss[1], ss[2]])
227 } else {
228 (0, [0, 0])
229 }
230 } else {
231 (0, [0, 0])
232 };
233 let b_stride = if is_affine {
234 let bn = proj_name.replace(".weight", ".biases");
235 if let Some((bdt, bs)) = storage.tensor_meta(&bn) {
236 bs[1] * bs[2] * bdt.size_in_bytes()
237 } else {
238 0
239 }
240 } else {
241 0
242 };
243 StackedProjMeta {
244 weight_byte_stride: w_stride,
245 weight_shape: [wshape[1], wshape[2]],
246 scales_byte_stride: s_stride,
247 scales_shape: s_shape,
248 biases_byte_stride: b_stride,
249 }

Callers

nothing calls this directly

Calls 3

get_expert_uncachedMethod · 0.80
tensor_metaMethod · 0.45
has_tensorMethod · 0.45

Tested by

no test coverage detected