use std::collections::HashMap;
use std::path::Path;
use crate::autograd::Tensor;
use crate::format::v2::AprV2Reader;
use crate::models::bert::config::BertConfig;
use crate::models::bert::load::{detect_bert_prefix, read_tensor};
use super::error::SetFitError;
use super::tokenizer::sha256_hex;
pub const PINNED_REVISION: &str = "1110a243fdf4706b3f48f1d95db1a4f5529b4d41";
pub const PINNED_TOKENIZER_SHA256: &str =
"be50c3628f2bf5bb5e3a7f17b1f74611b2561a3a27eeab05e5aa30f411572037";
pub const PINNED_MAX_SEQ_LENGTH: usize = 256;
pub const PINNED_ACTIVATION: &str = "gelu";
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct VocabRemap {
orig_to_slice: HashMap<u32, u32>,
slice_to_orig: Vec<u32>,
}
impl VocabRemap {
#[cfg(feature = "conformance-fixtures")]
pub(crate) fn from_json_bytes(bytes: &[u8], slice_vocab: usize) -> Result<Self, SetFitError> {
let wire: VocabRemapWire =
serde_json::from_slice(bytes).map_err(|e| SetFitError::RemapInvalid {
reason: format!("vocab_remap.json is not parseable: {e}"),
})?;
if wire.slice_to_orig.len() != slice_vocab {
return Err(SetFitError::RemapInvalid {
reason: format!(
"slice_to_orig has {} entries but the slice vocabulary is {slice_vocab}",
wire.slice_to_orig.len()
),
});
}
if wire.orig_to_slice.len() != slice_vocab {
return Err(SetFitError::RemapInvalid {
reason: format!(
"orig_to_slice has {} entries but the slice vocabulary is {slice_vocab}",
wire.orig_to_slice.len()
),
});
}
let bound = u32::try_from(slice_vocab).map_err(|_| SetFitError::RemapInvalid {
reason: format!("slice vocabulary {slice_vocab} does not fit in u32"),
})?;
for (canonical, slice_row) in &wire.orig_to_slice {
if *slice_row >= bound {
return Err(SetFitError::RemapInvalid {
reason: format!(
"canonical id {canonical} maps to slice row {slice_row}, \
which is outside a {slice_vocab}-row table"
),
});
}
}
for (canonical, slice_row) in &wire.orig_to_slice {
let back = wire.slice_to_orig[*slice_row as usize];
if back != *canonical {
return Err(SetFitError::RemapInvalid {
reason: format!(
"orig_to_slice[{canonical}] = {slice_row} but slice_to_orig[{slice_row}] \
= {back}; the two directions disagree"
),
});
}
}
for (row, canonical) in wire.slice_to_orig.iter().enumerate() {
match wire.orig_to_slice.get(canonical) {
Some(back) if usize::try_from(*back).ok() == Some(row) => {}
Some(back) => {
return Err(SetFitError::RemapInvalid {
reason: format!(
"slice_to_orig[{row}] = {canonical} but orig_to_slice[{canonical}] \
= {back}; the two directions disagree"
),
})
}
None => {
return Err(SetFitError::RemapInvalid {
reason: format!(
"slice_to_orig[{row}] = {canonical} has no orig_to_slice entry; \
the two directions disagree"
),
})
}
}
}
Ok(Self {
orig_to_slice: wire.orig_to_slice,
slice_to_orig: wire.slice_to_orig,
})
}
pub fn to_slice_row(&self, canonical: u32) -> Result<u32, SetFitError> {
self.orig_to_slice
.get(&canonical)
.copied()
.ok_or(SetFitError::VocabOutOfSlice {
canonical_id: canonical,
})
}
pub(crate) fn from_slice_to_orig(slice_to_orig: Vec<u32>) -> Result<Self, SetFitError> {
if slice_to_orig.is_empty() {
return Err(SetFitError::RemapInvalid {
reason: "slice_to_orig is empty; a slice with no vocabulary cannot gather"
.to_string(),
});
}
u32::try_from(slice_to_orig.len()).map_err(|_| SetFitError::RemapInvalid {
reason: format!(
"slice vocabulary {} does not fit in u32",
slice_to_orig.len()
),
})?;
let mut orig_to_slice: HashMap<u32, u32> = HashMap::with_capacity(slice_to_orig.len());
for (row, canonical) in slice_to_orig.iter().enumerate() {
let row_u32 = u32::try_from(row).map_err(|_| SetFitError::RemapInvalid {
reason: format!("slice row {row} does not fit in u32"),
})?;
if let Some(first) = orig_to_slice.insert(*canonical, row_u32) {
return Err(SetFitError::RemapInvalid {
reason: format!(
"canonical id {canonical} appears at slice rows {first} and {row}; \
the map is not injective and two tokens would share one embedding row"
),
});
}
}
Ok(Self {
orig_to_slice,
slice_to_orig,
})
}
#[must_use]
pub fn slice_to_orig(&self) -> &[u32] {
&self.slice_to_orig
}
#[must_use]
pub fn slice_vocab(&self) -> usize {
self.slice_to_orig.len()
}
#[must_use]
pub fn to_canonical(&self, slice_row: u32) -> Option<u32> {
self.slice_to_orig.get(slice_row as usize).copied()
}
}
#[cfg(feature = "conformance-fixtures")]
#[derive(serde::Deserialize)]
struct VocabRemapWire {
orig_to_slice: HashMap<u32, u32>,
slice_to_orig: Vec<u32>,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ModelDims {
pub hidden: usize,
pub layers: usize,
pub heads: usize,
pub intermediate: usize,
pub vocab: usize,
pub max_positions: usize,
pub type_vocab: usize,
pub pad_token_id: u32,
}
#[derive(Debug, Clone, serde::Deserialize)]
pub struct SliceConfig {
pub hidden: usize,
pub heads: usize,
pub head_dim: usize,
pub num_layers: usize,
pub intermediate: usize,
pub vocab: usize,
pub positions: usize,
pub type_vocab_size: usize,
pub layer_norm_eps: f64,
pub pad_token_id: u32,
pub hidden_act: String,
pub source_revision: String,
pub tokenizer_sha256: String,
}
impl SliceConfig {
#[cfg(feature = "conformance-fixtures")]
pub(crate) fn from_json_bytes(bytes: &[u8]) -> Result<Self, SetFitError> {
serde_json::from_slice(bytes).map_err(|e| SetFitError::ImportIo {
path: "slice_config.json".to_string(),
reason: e.to_string(),
})
}
}
#[derive(Debug, Clone, serde::Deserialize)]
struct HfBertConfig {
architectures: Vec<String>,
attention_probs_dropout_prob: f64,
hidden_act: String,
hidden_dropout_prob: f64,
hidden_size: usize,
intermediate_size: usize,
layer_norm_eps: f64,
max_position_embeddings: usize,
model_type: String,
num_attention_heads: usize,
num_hidden_layers: usize,
pad_token_id: u32,
position_embedding_type: String,
type_vocab_size: usize,
vocab_size: usize,
}
#[derive(Debug, Clone, serde::Deserialize)]
struct SentenceTransformerModule {
#[allow(dead_code)]
idx: usize,
#[allow(dead_code)]
path: String,
#[serde(rename = "type")]
kind: String,
}
#[derive(Debug, Clone, serde::Deserialize)]
struct PoolingConfig {
word_embedding_dimension: usize,
pooling_mode_cls_token: bool,
pooling_mode_mean_tokens: bool,
pooling_mode_max_tokens: bool,
pooling_mode_mean_sqrt_len_tokens: bool,
}
#[derive(Debug, Clone, serde::Deserialize)]
struct SentenceBertConfig {
max_seq_length: usize,
}
pub struct MiniLmImport {
dims: ModelDims,
layer_norm_eps: f32,
reader: AprV2Reader,
tensor_prefix: &'static str,
revision: String,
tokenizer_sha256: String,
vocab_remap: Option<VocabRemap>,
}
impl std::fmt::Debug for MiniLmImport {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MiniLmImport")
.field("dims", &self.dims)
.field("revision", &self.revision)
.field("tokenizer_sha256", &self.tokenizer_sha256)
.field("is_slice", &self.vocab_remap.is_some())
.finish()
}
}
impl MiniLmImport {
pub(crate) fn open(dir: &Path) -> Result<Self, SetFitError> {
let cfg = parse_hf_config(&read_required(dir, "config.json")?)?;
let (dims, layer_norm_eps) = validate_against_pin(&cfg)?;
validate_module_stack(&read_required(dir, "modules.json")?)?;
validate_pooling(&read_required(dir, "1_Pooling/config.json")?, dims.hidden)?;
if let Some(bytes) = read_optional(dir, "sentence_bert_config.json")? {
validate_sentence_bert(&bytes)?;
}
let tokenizer_bytes = read_required(dir, "tokenizer.json")?;
let got = sha256_hex(&tokenizer_bytes);
if got != PINNED_TOKENIZER_SHA256 {
return Err(SetFitError::TokenizerHashMismatch {
expected: PINNED_TOKENIZER_SHA256.to_string(),
got,
});
}
let (weights_name, weights) = read_weights(dir)?;
let reader = parse_apr(&weights, &weights_name)?;
let tensor_prefix = detect_bert_prefix(&reader);
load_and_check_tensors(&reader, tensor_prefix, &dims)?;
Ok(Self {
dims,
layer_norm_eps,
reader,
tensor_prefix,
revision: PINNED_REVISION.to_string(),
tokenizer_sha256: got,
vocab_remap: None,
})
}
#[cfg(feature = "conformance-fixtures")]
pub(crate) fn open_slice_fixture(
apr: &Path,
config: &SliceConfig,
remap: &VocabRemap,
) -> Result<Self, SetFitError> {
if config.hidden_act != PINNED_ACTIVATION {
return Err(SetFitError::UnsupportedActivation {
got: config.hidden_act.clone(),
});
}
if config.source_revision != PINNED_REVISION {
return Err(SetFitError::ImportConfigMismatch {
field: "source_revision".to_string(),
expected: PINNED_REVISION.to_string(),
got: config.source_revision.clone(),
});
}
if config.tokenizer_sha256 != PINNED_TOKENIZER_SHA256 {
return Err(SetFitError::TokenizerHashMismatch {
expected: PINNED_TOKENIZER_SHA256.to_string(),
got: config.tokenizer_sha256.clone(),
});
}
if config.heads == 0 || config.head_dim == 0 || config.hidden == 0 {
return Err(SetFitError::ImportConfigMismatch {
field: "heads".to_string(),
expected: "non-zero heads, head_dim and hidden".to_string(),
got: format!(
"heads {} head_dim {} hidden {}",
config.heads, config.head_dim, config.hidden
),
});
}
if config.heads * config.head_dim != config.hidden {
return Err(SetFitError::ImportConfigMismatch {
field: "heads".to_string(),
expected: format!("hidden {} / head_dim {}", config.hidden, config.head_dim),
got: format!(
"heads {} (heads * head_dim = {})",
config.heads,
config.heads * config.head_dim
),
});
}
for (field, value) in [
("num_layers", config.num_layers),
("intermediate", config.intermediate),
("vocab", config.vocab),
("positions", config.positions),
("type_vocab_size", config.type_vocab_size),
] {
if value == 0 {
return Err(SetFitError::ImportConfigMismatch {
field: field.to_string(),
expected: "non-zero".to_string(),
got: "0".to_string(),
});
}
}
let layer_norm_eps = narrow_eps(config.layer_norm_eps)?;
if remap.slice_vocab() != config.vocab {
return Err(SetFitError::RemapInvalid {
reason: format!(
"remap covers {} rows but the slice vocabulary is {}",
remap.slice_vocab(),
config.vocab
),
});
}
let dims = ModelDims {
hidden: config.hidden,
layers: config.num_layers,
heads: config.heads,
intermediate: config.intermediate,
vocab: config.vocab,
max_positions: config.positions,
type_vocab: config.type_vocab_size,
pad_token_id: config.pad_token_id,
};
let bytes = std::fs::read(apr).map_err(|e| SetFitError::ImportIo {
path: apr.display().to_string(),
reason: e.to_string(),
})?;
let reader = parse_apr(&bytes, &apr.display().to_string())?;
let tensor_prefix = detect_bert_prefix(&reader);
load_and_check_tensors(&reader, tensor_prefix, &dims)?;
Ok(Self {
dims,
layer_norm_eps,
reader,
tensor_prefix,
revision: config.source_revision.clone(),
tokenizer_sha256: config.tokenizer_sha256.clone(),
vocab_remap: Some(remap.clone()),
})
}
#[must_use]
pub fn dims(&self) -> &ModelDims {
&self.dims
}
#[must_use]
pub fn layer_norm_eps(&self) -> f32 {
self.layer_norm_eps
}
#[must_use]
pub fn revision(&self) -> &str {
&self.revision
}
#[must_use]
pub fn tokenizer_sha256(&self) -> &str {
&self.tokenizer_sha256
}
#[must_use]
pub fn vocab_remap(&self) -> Option<&VocabRemap> {
self.vocab_remap.as_ref()
}
pub(crate) fn reader(&self) -> &AprV2Reader {
&self.reader
}
pub(crate) fn tensor_prefix(&self) -> &'static str {
self.tensor_prefix
}
}
pub(crate) const PINNED_HIDDEN_DROPOUT_PROB: f64 = 0.1;
pub(crate) const PINNED_ATTENTION_DROPOUT_PROB: f64 = 0.1;
const PINNED_POSITION_EMBEDDING_TYPE: &str = "absolute";
const PINNED_ARCHITECTURE: &str = "BertModel";
const PINNED_MODEL_TYPE: &str = "bert";
const WEIGHT_FILE_CANDIDATES: [&str; 2] = ["full_model.apr", "model.apr"];
pub(super) fn read_required(dir: &Path, name: &str) -> Result<Vec<u8>, SetFitError> {
std::fs::read(dir.join(name)).map_err(|e| SetFitError::ImportIo {
path: name.to_string(),
reason: e.to_string(),
})
}
fn read_optional(dir: &Path, name: &str) -> Result<Option<Vec<u8>>, SetFitError> {
let path = dir.join(name);
if !path.exists() {
return Ok(None);
}
read_required(dir, name).map(Some)
}
fn read_weights(dir: &Path) -> Result<(String, Vec<u8>), SetFitError> {
for name in WEIGHT_FILE_CANDIDATES {
if let Some(bytes) = read_optional(dir, name)? {
return Ok((name.to_string(), bytes));
}
}
Err(SetFitError::ImportIo {
path: WEIGHT_FILE_CANDIDATES[0].to_string(),
reason: format!(
"no APR weights present (looked for {})",
WEIGHT_FILE_CANDIDATES.join(", ")
),
})
}
fn parse_apr(bytes: &[u8], name: &str) -> Result<AprV2Reader, SetFitError> {
AprV2Reader::from_bytes(bytes).map_err(|e| SetFitError::ImportIo {
path: name.to_string(),
reason: format!("not a readable APR v2 container: {e}"),
})
}
fn parse_hf_config(bytes: &[u8]) -> Result<HfBertConfig, SetFitError> {
serde_json::from_slice(bytes).map_err(|e| SetFitError::ImportIo {
path: "config.json".to_string(),
reason: e.to_string(),
})
}
fn mismatch(
field: &str,
expected: impl std::fmt::Display,
got: impl std::fmt::Display,
) -> SetFitError {
SetFitError::ImportConfigMismatch {
field: field.to_string(),
expected: expected.to_string(),
got: got.to_string(),
}
}
fn check_usize(field: &str, got: usize, expected: usize) -> Result<(), SetFitError> {
if got == expected {
Ok(())
} else {
Err(mismatch(field, expected, got))
}
}
fn narrow_eps(value: f64) -> Result<f32, SetFitError> {
#[allow(clippy::cast_possible_truncation)]
let narrowed = value as f32;
if !narrowed.is_finite() || narrowed <= 0.0 {
return Err(mismatch("layer_norm_eps", "a finite positive value", value));
}
Ok(narrowed)
}
fn validate_against_pin(cfg: &HfBertConfig) -> Result<(ModelDims, f32), SetFitError> {
let pin = BertConfig::minilm_l6();
if cfg.architectures.len() != 1 || cfg.architectures[0] != PINNED_ARCHITECTURE {
return Err(SetFitError::UnsupportedArchitecture {
got: format!("{:?}", cfg.architectures),
});
}
if cfg.model_type != PINNED_MODEL_TYPE {
return Err(mismatch("model_type", PINNED_MODEL_TYPE, &cfg.model_type));
}
if cfg.hidden_act != PINNED_ACTIVATION {
return Err(SetFitError::UnsupportedActivation {
got: cfg.hidden_act.clone(),
});
}
if cfg.position_embedding_type != PINNED_POSITION_EMBEDDING_TYPE {
return Err(mismatch(
"position_embedding_type",
PINNED_POSITION_EMBEDDING_TYPE,
&cfg.position_embedding_type,
));
}
check_usize("hidden_size", cfg.hidden_size, pin.hidden_dim)?;
check_usize("num_hidden_layers", cfg.num_hidden_layers, pin.num_layers)?;
check_usize(
"num_attention_heads",
cfg.num_attention_heads,
pin.num_heads,
)?;
check_usize(
"intermediate_size",
cfg.intermediate_size,
pin.intermediate_dim,
)?;
check_usize("vocab_size", cfg.vocab_size, pin.vocab_size)?;
check_usize(
"max_position_embeddings",
cfg.max_position_embeddings,
pin.max_position_embeddings,
)?;
check_usize("type_vocab_size", cfg.type_vocab_size, pin.type_vocab_size)?;
if cfg.pad_token_id != pin.pad_token_id {
return Err(mismatch("pad_token_id", pin.pad_token_id, cfg.pad_token_id));
}
let eps = narrow_eps(cfg.layer_norm_eps)?;
if eps != pin.layer_norm_eps {
return Err(mismatch(
"layer_norm_eps",
pin.layer_norm_eps,
cfg.layer_norm_eps,
));
}
if cfg.hidden_dropout_prob != PINNED_HIDDEN_DROPOUT_PROB {
return Err(mismatch(
"hidden_dropout_prob",
PINNED_HIDDEN_DROPOUT_PROB,
cfg.hidden_dropout_prob,
));
}
if cfg.attention_probs_dropout_prob != PINNED_ATTENTION_DROPOUT_PROB {
return Err(mismatch(
"attention_probs_dropout_prob",
PINNED_ATTENTION_DROPOUT_PROB,
cfg.attention_probs_dropout_prob,
));
}
Ok((
ModelDims {
hidden: cfg.hidden_size,
layers: cfg.num_hidden_layers,
heads: cfg.num_attention_heads,
intermediate: cfg.intermediate_size,
vocab: cfg.vocab_size,
max_positions: cfg.max_position_embeddings,
type_vocab: cfg.type_vocab_size,
pad_token_id: cfg.pad_token_id,
},
eps,
))
}
fn validate_module_stack(bytes: &[u8]) -> Result<(), SetFitError> {
let modules: Vec<SentenceTransformerModule> =
serde_json::from_slice(bytes).map_err(|e| SetFitError::ImportIo {
path: "modules.json".to_string(),
reason: e.to_string(),
})?;
let kinds: Vec<&str> = modules.iter().map(|m| m.kind.as_str()).collect();
let expected = [
"sentence_transformers.models.Transformer",
"sentence_transformers.models.Pooling",
"sentence_transformers.models.Normalize",
];
if kinds != expected {
return Err(SetFitError::UnsupportedPooling {
got: format!(
"modules.json declares {kinds:?}; the pin requires \
Transformer -> Pooling -> Normalize (the trailing Normalize is the \
normalize flag)"
),
});
}
Ok(())
}
fn validate_pooling(bytes: &[u8], hidden: usize) -> Result<(), SetFitError> {
let pooling: PoolingConfig =
serde_json::from_slice(bytes).map_err(|e| SetFitError::ImportIo {
path: "1_Pooling/config.json".to_string(),
reason: e.to_string(),
})?;
if !pooling.pooling_mode_mean_tokens {
return Err(SetFitError::UnsupportedPooling {
got: "pooling_mode_mean_tokens = false; the pin is mean pooling".to_string(),
});
}
for (field, enabled) in [
("pooling_mode_cls_token", pooling.pooling_mode_cls_token),
("pooling_mode_max_tokens", pooling.pooling_mode_max_tokens),
(
"pooling_mode_mean_sqrt_len_tokens",
pooling.pooling_mode_mean_sqrt_len_tokens,
),
] {
if enabled {
return Err(SetFitError::UnsupportedPooling {
got: format!("{field} = true; the pin is mean pooling ONLY"),
});
}
}
if pooling.word_embedding_dimension != hidden {
return Err(mismatch(
"word_embedding_dimension",
hidden,
pooling.word_embedding_dimension,
));
}
Ok(())
}
fn validate_sentence_bert(bytes: &[u8]) -> Result<(), SetFitError> {
let sbert: SentenceBertConfig =
serde_json::from_slice(bytes).map_err(|e| SetFitError::ImportIo {
path: "sentence_bert_config.json".to_string(),
reason: e.to_string(),
})?;
check_usize(
"max_seq_length",
sbert.max_seq_length,
PINNED_MAX_SEQ_LENGTH,
)
}
fn expected_tensor_specs(prefix: &str, dims: &ModelDims) -> Vec<(String, Vec<usize>)> {
let h = dims.hidden;
let im = dims.intermediate;
let mut specs = vec![
(
format!("{prefix}embeddings.word_embeddings.weight"),
vec![dims.vocab, h],
),
(
format!("{prefix}embeddings.position_embeddings.weight"),
vec![dims.max_positions, h],
),
(
format!("{prefix}embeddings.token_type_embeddings.weight"),
vec![dims.type_vocab, h],
),
(format!("{prefix}embeddings.LayerNorm.weight"), vec![h]),
(format!("{prefix}embeddings.LayerNorm.bias"), vec![h]),
];
for idx in 0..dims.layers {
let p = format!("{prefix}encoder.layer.{idx}");
for proj in ["query", "key", "value"] {
specs.push((format!("{p}.attention.self.{proj}.weight"), vec![h, h]));
specs.push((format!("{p}.attention.self.{proj}.bias"), vec![h]));
}
specs.push((format!("{p}.attention.output.dense.weight"), vec![h, h]));
specs.push((format!("{p}.attention.output.dense.bias"), vec![h]));
specs.push((format!("{p}.attention.output.LayerNorm.weight"), vec![h]));
specs.push((format!("{p}.attention.output.LayerNorm.bias"), vec![h]));
specs.push((format!("{p}.intermediate.dense.weight"), vec![im, h]));
specs.push((format!("{p}.intermediate.dense.bias"), vec![im]));
specs.push((format!("{p}.output.dense.weight"), vec![h, im]));
specs.push((format!("{p}.output.dense.bias"), vec![h]));
specs.push((format!("{p}.output.LayerNorm.weight"), vec![h]));
specs.push((format!("{p}.output.LayerNorm.bias"), vec![h]));
}
specs
}
fn load_and_check_tensors(
reader: &AprV2Reader,
prefix: &str,
dims: &ModelDims,
) -> Result<(), SetFitError> {
for (name, shape) in expected_tensor_specs(prefix, dims) {
let tensor: Tensor = read_tensor(reader, &name, &shape)?;
if let Some(position) = tensor.data().iter().position(|v| !v.is_finite()) {
return Err(SetFitError::NonFiniteTensor {
tensor: name,
position,
});
}
}
Ok(())
}
#[cfg(all(test, feature = "setfit"))]
#[path = "import_tests.rs"]
mod import_tests;