|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| use meganeura::Session;
|
| use meganeura::data::safetensors::SafeTensorsModel;
|
|
|
| use crate::dinov3::Config;
|
| use crate::preprocess::conv_weight_to_matmul;
|
|
|
| type Error = Box<dyn std::error::Error>;
|
|
|
|
|
|
|
|
|
|
|
| fn layer_prefix(model: &SafeTensorsModel) -> Result<&'static str, Error> {
|
| let info = model.tensor_info();
|
| for candidate in ["layer.", "model.layer."] {
|
| if info.keys().any(|k| k.starts_with(candidate)) {
|
| return Ok(candidate);
|
| }
|
| }
|
| Err(format!(
|
| "no encoder layers found; checkpoint has {} tensors, e.g. {:?}",
|
| info.len(),
|
| info.keys().take(5).collect::<Vec<_>>()
|
| )
|
| .into())
|
| }
|
|
|
|
|
| fn embedding_tensor(model: &SafeTensorsModel, name: &str) -> Result<Vec<f32>, Error> {
|
| if model.tensor_info().contains_key(name) {
|
| model.tensor_f32_auto(name)
|
| } else {
|
| model.tensor_f32_auto(&format!("model.{name}"))
|
| }
|
| }
|
|
|
|
|
|
|
|
|
|
|
|
|
| pub fn load_encoder(
|
| session: &mut Session,
|
| model: &SafeTensorsModel,
|
| config: &Config,
|
| ) -> Result<(), Error> {
|
| let hidden = config.hidden_size;
|
| let prefix = layer_prefix(model)?;
|
| log::info!("checkpoint layer prefix: {prefix:?}");
|
|
|
|
|
| let conv = embedding_tensor(model, "embeddings.patch_embeddings.weight")?;
|
| let expected = hidden * config.patch_dim();
|
| if conv.len() != expected {
|
| return Err(format!(
|
| "patch embedding has {} values, expected {expected} \
|
| (is the checkpoint's patch_size or hidden_size different?)",
|
| conv.len()
|
| )
|
| .into());
|
| }
|
| session.set_parameter(
|
| "embeddings.patch_embeddings.weight",
|
| &conv_weight_to_matmul(&conv, hidden, config.patch_dim()),
|
| );
|
| session.set_parameter(
|
| "embeddings.patch_embeddings.bias",
|
| &embedding_tensor(model, "embeddings.patch_embeddings.bias")?,
|
| );
|
|
|
|
|
|
|
|
|
|
|
|
|
| let cls = embedding_tensor(model, "embeddings.cls_token")?;
|
| let registers = embedding_tensor(model, "embeddings.register_tokens")?;
|
| if cls.len() != hidden {
|
| return Err(format!("cls_token has {} values, expected {hidden}", cls.len()).into());
|
| }
|
| if registers.len() != config.num_register_tokens * hidden {
|
| return Err(format!(
|
| "register_tokens has {} values, expected {} * {hidden}",
|
| registers.len(),
|
| config.num_register_tokens
|
| )
|
| .into());
|
| }
|
| let mut prefix_tokens = Vec::with_capacity(config.num_prefix_tokens() * hidden);
|
| prefix_tokens.extend_from_slice(&cls);
|
| prefix_tokens.extend_from_slice(®isters);
|
| session.set_parameter("prefix_tokens", &prefix_tokens);
|
|
|
|
|
| for i in 0..config.num_hidden_layers {
|
| let src = format!("{prefix}{i}");
|
| let dst = format!("layer.{i}");
|
|
|
| for norm in ["norm1", "norm2"] {
|
| for part in ["weight", "bias"] {
|
| session.set_parameter(
|
| &format!("{dst}.{norm}.{part}"),
|
| &model.tensor_f32_auto(&format!("{src}.{norm}.{part}"))?,
|
| );
|
| }
|
| }
|
|
|
|
|
| for (proj, has_bias) in [
|
| ("q_proj", true),
|
| ("k_proj", false),
|
| ("v_proj", true),
|
| ("o_proj", true),
|
| ] {
|
| session.set_parameter(
|
| &format!("{dst}.attention.{proj}.weight"),
|
| &model.tensor_f32_auto_transposed(&format!("{src}.attention.{proj}.weight"))?,
|
| );
|
| if has_bias {
|
| session.set_parameter(
|
| &format!("{dst}.attention.{proj}.bias"),
|
| &model.tensor_f32_auto(&format!("{src}.attention.{proj}.bias"))?,
|
| );
|
| }
|
| }
|
|
|
| for proj in ["up_proj", "down_proj"] {
|
| session.set_parameter(
|
| &format!("{dst}.mlp.{proj}.weight"),
|
| &model.tensor_f32_auto_transposed(&format!("{src}.mlp.{proj}.weight"))?,
|
| );
|
| session.set_parameter(
|
| &format!("{dst}.mlp.{proj}.bias"),
|
| &model.tensor_f32_auto(&format!("{src}.mlp.{proj}.bias"))?,
|
| );
|
| }
|
|
|
| for ls in ["layer_scale1", "layer_scale2"] {
|
| session.set_parameter(
|
| &format!("{dst}.{ls}.lambda1"),
|
| &model.tensor_f32_auto(&format!("{src}.{ls}.lambda1"))?,
|
| );
|
| }
|
| }
|
|
|
|
|
| for part in ["weight", "bias"] {
|
| session.set_parameter(
|
| &format!("norm.{part}"),
|
| &embedding_tensor(model, &format!("norm.{part}"))?,
|
| );
|
| }
|
|
|
| log::info!(
|
| "loaded DINOv3 weights: {} layers, hidden {hidden}",
|
| config.num_hidden_layers
|
| );
|
| Ok(())
|
| }
|
|
|