spark_model/weight_loader/
laguna.rs

1// SPDX-License-Identifier: AGPL-3.0-only
2
3//! Poolside Laguna-S-2.1 weight loader (target model only; DFlash is separate).
4
5mod load_layers;
6
7use anyhow::Result;
8use atlas_core::config::ModelConfig;
9use spark_runtime::gpu::GpuBackend;
10use spark_runtime::kv_cache::KvCacheDtype;
11use spark_runtime::weights::WeightStore;
12
13use super::ModelWeightLoader;
14use crate::layer::TransformerLayer;
15use crate::weight_map::{DenseWeight, MtpWeights, dense};
16
17pub struct LagunaWeightLoader;
18
19impl ModelWeightLoader for LagunaWeightLoader {
20    fn supports_tp(&self) -> bool {
21        false
22    }
23
24    fn load_layers(
25        &self,
26        store: &WeightStore,
27        config: &ModelConfig,
28        gpu: &dyn GpuBackend,
29        layer_kv_dtypes: &[KvCacheDtype],
30    ) -> Result<Vec<Box<dyn TransformerLayer>>> {
31        load_layers::load_layers(store, config, gpu, layer_kv_dtypes)
32    }
33
34    fn load_embedding(
35        &self,
36        store: &WeightStore,
37        _config: &ModelConfig,
38        _gpu: &dyn GpuBackend,
39    ) -> Result<DenseWeight> {
40        dense(store, "model.embed_tokens.weight")
41    }
42
43    fn load_final_norm(
44        &self,
45        store: &WeightStore,
46        _config: &ModelConfig,
47        _gpu: &dyn GpuBackend,
48    ) -> Result<DenseWeight> {
49        dense(store, "model.norm.weight")
50    }
51
52    fn load_lm_head(
53        &self,
54        store: &WeightStore,
55        _config: &ModelConfig,
56        _gpu: &dyn GpuBackend,
57    ) -> Result<DenseWeight> {
58        dense(store, "lm_head.weight")
59    }
60
61    fn load_mtp_weights(
62        &self,
63        _store: &WeightStore,
64        _config: &ModelConfig,
65        _gpu: &dyn GpuBackend,
66    ) -> Result<Option<MtpWeights>> {
67        Ok(None)
68    }
69}