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

Method load

cake-core/src/models/common/text_model.rs:150–263  ·  view source on GitHub ↗

Load the shared model structure from the context. `default_eos_token` is the model-specific fallback EOS string. The type parameter `B` determines which block type to use for local layers.

(
        ctx: &mut Context,
        default_eos_token: &str,
    )

Source from the content-addressed store, hash-verified

148 /// `default_eos_token` is the model-specific fallback EOS string.
149 /// The type parameter `B` determines which block type to use for local layers.
150 pub async fn load<B: Forwarder + 'static>(
151 ctx: &mut Context,
152 default_eos_token: &str,
153 ) -> Result<Self> {
154 let config = ctx.config.as_ref().expect("No config specified");
155 let var_builder = ctx.var_builder.as_ref().expect("No var_builder specified");
156 let prefix = &config.model_prefix;
157
158 log::info!("loading embeddings (prefix={}) ...", prefix);
159 let embed_weight = var_builder
160 .pp(format!("{prefix}.embed_tokens"))
161 .get((config.vocab_size, config.hidden_size), "weight")?;
162
163 log::info!("loading lm_head ...");
164 let lm_head_weight = if config.tie_word_embeddings {
165 log::info!(" using tied word embeddings (lm_head = embed_tokens)");
166 embed_weight.clone()
167 } else {
168 // Try multiple lm_head locations:
169 // 1. Root: lm_head.weight (LLaMA, Qwen2)
170 // 2. Prefixed: {prefix}.lm_head.weight (Qwen3.5)
171 // 3. Parent: {parent}.lm_head.weight (ConditionalGeneration models where
172 // prefix="language_model.model" but lm_head is at "language_model.lm_head")
173 let lm_head_shape = (config.vocab_size, config.hidden_size);
174 var_builder.pp("lm_head").get(lm_head_shape, "weight")
175 .or_else(|_| var_builder.pp(format!("{prefix}.lm_head")).get(lm_head_shape, "weight"))
176 .or_else(|_| {
177 // Try parent prefix (strip last component: "a.b.c" -> "a.b")
178 if let Some(parent) = prefix.rsplit_once('.').map(|(p, _)| p) {
179 var_builder.pp(format!("{parent}.lm_head")).get(lm_head_shape, "weight")
180 } else {
181 Err(candle_core::Error::Msg(format!("cannot find lm_head.weight (tried root, {prefix}.lm_head)")))
182 }
183 })?
184 };
185 let lm_head_weight = ctx.backend.preprocess_linear_weight(&lm_head_weight)?;
186
187 log::info!("loading {prefix}.norm ...");
188 let ln_f_weight = crate::models::common::load_rms_norm_weight(
189 config.hidden_size,
190 config.residual_rms_norm,
191 var_builder.pp(format!("{prefix}.norm")),
192 )?;
193 let ln_f_eps = config.rms_norm_eps as f32;
194
195 log::info!("loading {} blocks ...", config.num_hidden_layers);
196
197 // Two-pass loading: local layers first (no network wait), then remote
198 // layers (may block until workers finish loading). This overlaps
199 // master's local layer loading with worker startup time.
200 let mut blocks: Vec<Option<Box<dyn Forwarder>>> =
201 (0..config.num_hidden_layers).map(|_| None).collect();
202
203 // Pass 1: load local layers
204 for (i, block) in blocks.iter_mut().enumerate().take(config.num_hidden_layers) {
205 let block_layer_name = format!("{prefix}.layers.{i}");
206 if ctx.topology.get_node_for_layer(&block_layer_name).is_none() {
207 log::info!("loading {} ...", &block_layer_name);

Callers

nothing calls this directly

Calls 7

load_rms_norm_weightFunction · 0.85
load_tokenizerFunction · 0.85
create_logits_processorFunction · 0.85
get_node_for_layerMethod · 0.80
getMethod · 0.45
cloneMethod · 0.45

Tested by

no test coverage detected