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