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,
)
| 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 | } |
nothing calls this directly
no test coverage detected