pub mod artifact;
pub mod classify;
pub mod dropout_rng;
pub mod encoder;
pub mod error;
pub mod import;
pub mod loss;
pub mod tokenizer;
pub use artifact::{
artifact_sha256_hex, load_setfit_apr, read_setfit_apr_bytes_bounded, read_setfit_apr_parts,
write_setfit_apr, ProbeReplayDivergence, SetFitAprParts, SetFitArtifactDoc,
SetFitArtifactError, SetFitArtifactView, SetFitHeadDoc, SetFitPreprocessingDoc,
SetFitProbeRecord, VerifiedSetFitModel, MAX_ARTIFACT_BYTES, MAX_ENCODER_LAYERS,
NULLABLE_PATH_ALLOWLIST, PROBE_EMBEDDING_ABS_TOLERANCE, PROBE_LOGITS_ABS_TOLERANCE,
PROBE_PROBABILITIES_ABS_TOLERANCE, WALKED_SUBDOCUMENTS,
};
pub use classify::{
ClassifyError, ClassifyRequestDocument, ClassifyResponse, ClassifyResult,
CLASSIFY_SCHEMA_VERSION, MAX_BATCH_TEXTS, MAX_REQUEST_BODY_BYTES,
PROBABILITY_MASS_ABS_TOLERANCE,
};
pub use dropout_rng::{DropoutRngError, SiteDropout};
pub use encoder::{
BertSentenceEncoder, ExecutionBackend, L2_EPS, NORMALIZATION_POLICY, POOLING_POLICY,
};
pub use error::SetFitError;
pub use import::{
MiniLmImport, ModelDims, SliceConfig, VocabRemap, PINNED_ACTIVATION, PINNED_MAX_SEQ_LENGTH,
PINNED_REVISION, PINNED_TOKENIZER_SHA256,
};
pub use loss::pair_cosine_mse;
pub use tokenizer::{
InputProvenance, MiniLmTokenizer, SentenceBatch, TruncationFact, MAX_SEQUENCE_LENGTH,
PADDING_MODE,
};
use std::collections::BTreeMap;
use std::path::Path;
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
#[serde(deny_unknown_fields)]
pub struct EncoderArchitecture {
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,
pub vocab_remap: Option<Vec<u32>>,
}
use crate::autograd::Tensor;
use crate::nn::Module;
#[derive(Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, Hash)]
pub enum FreezeGroup {
Embeddings,
LayerAttention(usize),
LayerFfn(usize),
LayerNorm(usize),
}
impl FreezeGroup {
#[must_use]
pub fn layer(self) -> Option<usize> {
match self {
Self::Embeddings => None,
Self::LayerAttention(n) | Self::LayerFfn(n) | Self::LayerNorm(n) => Some(n),
}
}
#[must_use]
pub fn name_prefixes(self) -> Vec<String> {
match self {
Self::Embeddings => vec!["embeddings.".to_string()],
Self::LayerAttention(n) => vec![
format!("encoder.layer.{n}.attention.self."),
format!("encoder.layer.{n}.attention.output.dense."),
],
Self::LayerFfn(n) => vec![
format!("encoder.layer.{n}.intermediate."),
format!("encoder.layer.{n}.output.dense."),
],
Self::LayerNorm(n) => vec![
format!("encoder.layer.{n}.attention.output.LayerNorm."),
format!("encoder.layer.{n}.output.LayerNorm."),
],
}
}
#[must_use]
pub fn matches(self, name: &str) -> bool {
self.name_prefixes().iter().any(|p| name.starts_with(p))
}
}
pub struct SetFitMiniLm {
tokenizer: MiniLmTokenizer,
encoder: BertSentenceEncoder,
freeze: Vec<FreezeGroup>,
}
impl std::fmt::Debug for SetFitMiniLm {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("SetFitMiniLm")
.field("encoder", &self.encoder)
.field("freeze", &self.freeze)
.finish_non_exhaustive()
}
}
impl SetFitMiniLm {
pub fn from_pretrained_dir(dir: &Path, root_seed: u64) -> Result<Self, SetFitError> {
let tokenizer =
MiniLmTokenizer::from_bytes(&import::read_required(dir, "tokenizer.json")?)?;
let import = MiniLmImport::open(dir)?;
let encoder = BertSentenceEncoder::from_import(&import, root_seed)?;
Ok(Self {
tokenizer,
encoder,
freeze: Vec::new(),
})
}
#[cfg(feature = "conformance-fixtures")]
pub fn from_slice_fixture(fixture_dir: &Path, root_seed: u64) -> Result<Self, SetFitError> {
let tokenizer =
MiniLmTokenizer::from_bytes(&import::read_required(fixture_dir, "tokenizer.json")?)?;
let config = SliceConfig::from_json_bytes(&import::read_required(
fixture_dir,
"slice_config.json",
)?)?;
let remap = VocabRemap::from_json_bytes(
&import::read_required(fixture_dir, "vocab_remap.json")?,
config.vocab,
)?;
let import = MiniLmImport::open_slice_fixture(
&fixture_dir.join("slice_model.apr"),
&config,
&remap,
)?;
let encoder = BertSentenceEncoder::from_import(&import, root_seed)?;
Ok(Self {
tokenizer,
encoder,
freeze: Vec::new(),
})
}
pub fn from_bundle_parts(
tokenizer_bytes: &[u8],
arch: &EncoderArchitecture,
tensors: BTreeMap<String, (Vec<usize>, Vec<f32>)>,
root_seed: u64,
) -> Result<Self, SetFitError> {
let observed = tokenizer::sha256_hex(tokenizer_bytes);
if observed != arch.tokenizer_sha256 {
return Err(SetFitError::TokenizerHashMismatch {
expected: arch.tokenizer_sha256.clone(),
got: observed,
});
}
let tokenizer = MiniLmTokenizer::from_bytes(tokenizer_bytes)?;
let encoder = BertSentenceEncoder::from_named_tensors(arch, tensors, root_seed)?;
Ok(Self {
tokenizer,
encoder,
freeze: Vec::new(),
})
}
#[must_use]
pub fn tokenizer_sha256(&self) -> &str {
self.tokenizer.tokenizer_sha256()
}
#[must_use]
pub fn tokenizer_bytes(&self) -> &[u8] {
self.tokenizer.source_bytes()
}
#[must_use]
pub fn architecture(&self) -> EncoderArchitecture {
let dims = self.encoder.dims();
EncoderArchitecture {
hidden: dims.hidden,
heads: dims.heads,
head_dim: dims.hidden / dims.heads.max(1),
num_layers: dims.layers,
intermediate: dims.intermediate,
vocab: dims.vocab,
positions: dims.max_positions,
type_vocab_size: dims.type_vocab,
layer_norm_eps: f64::from(self.encoder.layer_norm_eps()),
pad_token_id: dims.pad_token_id,
hidden_act: PINNED_ACTIVATION.to_string(),
source_revision: self.encoder.source_revision().to_string(),
tokenizer_sha256: self.tokenizer.tokenizer_sha256().to_string(),
vocab_remap: self
.encoder
.vocab_remap()
.map(|remap| remap.slice_to_orig().to_vec()),
}
}
#[must_use]
pub fn named_parameters(&self) -> Vec<(String, &Tensor)> {
self.encoder.named_parameters()
}
#[must_use]
pub fn architecture_fingerprint(&self) -> String {
self.encoder.architecture_fingerprint()
}
#[must_use]
pub fn num_layers(&self) -> usize {
self.encoder.num_layers()
}
#[must_use]
pub fn root_seed(&self) -> u64 {
self.encoder.root_seed()
}
#[must_use]
pub fn training(&self) -> bool {
self.encoder.training()
}
pub fn encode_texts(&self, texts: &[&str]) -> Result<Tensor, SetFitError> {
let batch = self.tokenizer.encode_batch(texts)?;
self.encoder.encode(&batch)
}
pub fn encode_texts_traced(
&self,
texts: &[&str],
) -> Result<(Tensor, ExecutionBackend), SetFitError> {
let batch = self.tokenize_batch(texts)?;
self.encode_batch_traced(&batch)
}
pub(crate) fn tokenize_batch(&self, texts: &[&str]) -> Result<SentenceBatch, SetFitError> {
self.tokenizer.encode_batch(texts)
}
pub(crate) fn encode_batch_traced(
&self,
batch: &SentenceBatch,
) -> Result<(Tensor, ExecutionBackend), SetFitError> {
self.encoder.encode_with_backend(batch)
}
pub fn set_training(&mut self, training: bool) {
self.encoder.set_training(training);
}
pub fn set_forward_ordinal(&mut self, forward_ordinal: u64) -> Result<(), SetFitError> {
self.encoder.set_forward_ordinal(forward_ordinal)
}
#[must_use]
pub fn forward_ordinal(&self) -> u64 {
self.encoder.forward_ordinal()
}
#[cfg(feature = "conformance-fixtures")]
#[must_use]
pub fn encoder(&self) -> &BertSentenceEncoder {
&self.encoder
}
#[cfg(feature = "conformance-fixtures")]
pub fn tokenize(&self, texts: &[&str]) -> Result<SentenceBatch, SetFitError> {
self.tokenizer.encode_batch(texts)
}
pub fn apply_freeze(&mut self, groups: &[FreezeGroup]) -> Result<(), SetFitError> {
let layers = self.encoder.num_layers();
let all: Vec<String> = self
.encoder
.named_parameters()
.into_iter()
.map(|(n, _)| n)
.collect();
for g in groups {
if let Some(n) = g.layer() {
if n >= layers {
return Err(SetFitError::FreezeGroupInvalid {
reason: format!(
"{g:?} names layer {n}, but this encoder has {layers} layers \
(valid indices 0..{})",
layers.saturating_sub(1)
),
});
}
}
if !all.iter().any(|name| g.matches(name)) {
return Err(SetFitError::FreezeGroupInvalid {
reason: format!(
"{g:?} addresses ZERO named parameters (prefixes {:?}); the encoder's \
parameter naming has drifted from the freeze mapping",
g.name_prefixes()
),
});
}
}
let mut normalized = groups.to_vec();
normalized.sort_unstable();
normalized.dedup();
self.set_all_requires_grad(true);
self.freeze = normalized;
let frozen = self.frozen_names();
for (name, t) in self.encoder.named_parameters_mut() {
if frozen.contains(&name) {
t.requires_grad_(false);
}
}
Ok(())
}
pub fn clear_freeze(&mut self) {
self.freeze.clear();
self.set_all_requires_grad(true);
}
fn frozen_names(&self) -> Vec<String> {
self.encoder
.named_parameters()
.into_iter()
.map(|(n, _)| n)
.filter(|n| self.freeze.iter().any(|g| g.matches(n)))
.collect()
}
fn set_all_requires_grad(&mut self, requires: bool) {
for (_, t) in self.encoder.named_parameters_mut() {
t.requires_grad_(requires);
}
}
#[must_use]
pub fn freeze_policy(&self) -> Vec<FreezeGroup> {
self.freeze.clone()
}
#[must_use]
pub fn trainable_parameters_mut(&mut self) -> Vec<(String, &mut Tensor)> {
let frozen = self.frozen_names();
self.encoder
.named_parameters_mut()
.into_iter()
.filter(|(n, _)| !frozen.contains(n))
.collect()
}
#[must_use]
pub fn frozen_parameters(&self) -> Vec<(String, &Tensor)> {
self.encoder
.named_parameters()
.into_iter()
.filter(|(n, _)| self.freeze.iter().any(|g| g.matches(n)))
.collect()
}
}
#[cfg(all(test, feature = "setfit"))]
#[path = "model_tests.rs"]
mod model_tests;