use std::collections::BTreeMap;
use std::sync::Arc;
use crate::autograd::{
additive_attention_mask, embedding_gather, l2_normalize_rows, masked_mean_pool, OpError, Tensor,
};
use crate::models::bert::load::{read_tensor, BertLoadError};
use crate::nn::transformer::AttentionDropoutMasks;
use crate::nn::{LayerNorm, Linear, Module, MultiHeadAttention};
use super::dropout_rng::{self, SiteDropout};
use super::error::SetFitError;
use super::import::{MiniLmImport, ModelDims, VocabRemap, PINNED_ACTIVATION};
use super::tokenizer::{SentenceBatch, MAX_SEQUENCE_LENGTH};
use super::EncoderArchitecture;
pub const L2_EPS: f32 = 1e-12;
pub const POOLING_POLICY: &str = "masked_mean";
pub const NORMALIZATION_POLICY: &str = "l2";
const DROPOUT_P: f32 = 0.1;
const EMBEDDINGS_DROPOUT_SITE: &str = "embeddings.dropout";
fn attention_probs_site(layer: usize) -> String {
format!("encoder.layer.{layer}.attention.self.dropout")
}
fn attention_output_site(layer: usize) -> String {
format!("encoder.layer.{layer}.attention.output.dropout")
}
fn ffn_output_site(layer: usize) -> String {
format!("encoder.layer.{layer}.output.dropout")
}
fn site_dropout(root_seed: u64, site: &str, p: f32) -> Result<Arc<SiteDropout>, SetFitError> {
Ok(Arc::new(SiteDropout::new(root_seed, site, p)?))
}
struct EncoderLayer {
attention: MultiHeadAttention,
attention_probs_dropout: Arc<SiteDropout>,
attention_output_dropout: Arc<SiteDropout>,
attention_layer_norm: LayerNorm,
intermediate: Linear,
output_dense: Linear,
output_dropout: Arc<SiteDropout>,
output_layer_norm: LayerNorm,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ExecutionBackend {
device: &'static str,
kernel: &'static str,
}
impl ExecutionBackend {
#[must_use]
pub fn identity(&self) -> String {
format!("{}:setfit-core:{}", self.device, self.kernel)
}
#[must_use]
pub fn device(&self) -> &'static str {
self.device
}
#[must_use]
pub fn kernel(&self) -> &'static str {
self.kernel
}
}
pub struct BertSentenceEncoder {
word_embeddings: Tensor,
position_embeddings: Tensor,
token_type_embeddings: Tensor,
embeddings_layer_norm: LayerNorm,
embeddings_dropout: Arc<SiteDropout>,
layers: Vec<EncoderLayer>,
dims: ModelDims,
remap: Option<VocabRemap>,
layer_norm_eps: f32,
source_revision: String,
tokenizer_sha256: String,
root_seed: u64,
training: bool,
forward_ordinal: u64,
}
impl BertSentenceEncoder {
#[must_use]
pub fn architecture_fingerprint(&self) -> String {
format!(
"minilm-slice-h{}-l{}-a{}-i{}-v{}",
self.dims.hidden,
self.dims.layers,
self.dims.heads,
self.dims.intermediate,
self.dims.vocab,
)
}
}
impl std::fmt::Debug for BertSentenceEncoder {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BertSentenceEncoder")
.field("dims", &self.dims)
.field("max_seq", &self.max_seq())
.field("is_slice", &self.remap.is_some())
.field("training", &self.training)
.finish_non_exhaustive()
}
}
impl BertSentenceEncoder {
pub(crate) fn from_import(import: &MiniLmImport, root_seed: u64) -> Result<Self, SetFitError> {
let reader = import.reader();
let prefix = import.tensor_prefix();
Self::assemble(
import.dims().clone(),
import.layer_norm_eps(),
import.vocab_remap().cloned(),
import.revision().to_string(),
import.tokenizer_sha256().to_string(),
root_seed,
&|name: &str, shape: &[usize]| -> Result<Tensor, SetFitError> {
Ok(read_tensor(reader, &format!("{prefix}{name}"), shape)?.requires_grad())
},
)
}
pub(crate) fn from_named_tensors(
arch: &EncoderArchitecture,
tensors: BTreeMap<String, (Vec<usize>, Vec<f32>)>,
root_seed: u64,
) -> Result<Self, SetFitError> {
if arch.hidden_act != PINNED_ACTIVATION {
return Err(SetFitError::UnsupportedActivation {
got: arch.hidden_act.clone(),
});
}
if arch.heads == 0 || arch.head_dim.checked_mul(arch.heads) != Some(arch.hidden) {
return Err(SetFitError::ImportConfigMismatch {
field: "head_dim".to_string(),
expected: format!("hidden / heads = {} / {}", arch.hidden, arch.heads),
got: arch.head_dim.to_string(),
});
}
let supplied_elements: usize = tensors.values().map(|(_, data)| data.len()).sum();
for (field, value, ceiling, ceiling_of) in [
(
"num_layers",
arch.num_layers,
tensors.len(),
"supplied tensors",
),
(
"hidden",
arch.hidden,
supplied_elements,
"supplied elements",
),
(
"intermediate",
arch.intermediate,
supplied_elements,
"supplied elements",
),
("vocab", arch.vocab, supplied_elements, "supplied elements"),
(
"positions",
arch.positions,
supplied_elements,
"supplied elements",
),
(
"type_vocab_size",
arch.type_vocab_size,
supplied_elements,
"supplied elements",
),
] {
if value > ceiling {
return Err(SetFitError::ImportConfigMismatch {
field: field.to_string(),
expected: format!("at most {ceiling} ({ceiling_of})"),
got: value.to_string(),
});
}
}
let dims = ModelDims {
hidden: arch.hidden,
layers: arch.num_layers,
heads: arch.heads,
intermediate: arch.intermediate,
vocab: arch.vocab,
max_positions: arch.positions,
type_vocab: arch.type_vocab_size,
pad_token_id: arch.pad_token_id,
};
let remap = match &arch.vocab_remap {
Some(slice_to_orig) => {
let remap = VocabRemap::from_slice_to_orig(slice_to_orig.clone())?;
if remap.slice_vocab() != dims.vocab {
return Err(SetFitError::RemapInvalid {
reason: format!(
"the remap has {} rows but the embedding table has {}",
remap.slice_vocab(),
dims.vocab
),
});
}
Some(remap)
}
None => None,
};
let remaining = core::cell::RefCell::new(tensors);
let encoder = Self::assemble(
dims,
arch.layer_norm_eps as f32,
remap,
arch.source_revision.clone(),
arch.tokenizer_sha256.clone(),
root_seed,
&|name: &str, shape: &[usize]| -> Result<Tensor, SetFitError> {
let (got_shape, data) = remaining.borrow_mut().remove(name).ok_or_else(|| {
SetFitError::ImportTensor(BertLoadError {
tensor: name.to_string(),
reason: "tensor not present in the bundle".to_string(),
})
})?;
if got_shape.as_slice() != shape {
return Err(SetFitError::ImportTensor(BertLoadError {
tensor: name.to_string(),
reason: format!("shape mismatch: got {got_shape:?}, expected {shape:?}"),
}));
}
let expected: usize = shape.iter().product();
if data.len() != expected {
return Err(SetFitError::ImportTensor(BertLoadError {
tensor: name.to_string(),
reason: format!(
"element count mismatch: got {}, expected {expected} (shape {shape:?})",
data.len()
),
}));
}
Ok(Tensor::from_vec(data, shape).requires_grad())
},
)?;
let leftover = remaining.into_inner();
if let Some((name, _)) = leftover.into_iter().next() {
return Err(SetFitError::ImportTensor(BertLoadError {
tensor: name,
reason: "the bundle carries a tensor this architecture does not name".to_string(),
}));
}
Ok(encoder)
}
fn assemble(
dims: ModelDims,
eps: f32,
remap: Option<VocabRemap>,
source_revision: String,
tokenizer_sha256: String,
root_seed: u64,
read: &dyn Fn(&str, &[usize]) -> Result<Tensor, SetFitError>,
) -> Result<Self, SetFitError> {
let h = dims.hidden;
let mut embeddings_layer_norm = LayerNorm::with_eps(&[h], eps);
embeddings_layer_norm.set_weight(read("embeddings.LayerNorm.weight", &[h])?);
embeddings_layer_norm.set_bias(read("embeddings.LayerNorm.bias", &[h])?);
let mut layers = Vec::with_capacity(dims.layers);
for i in 0..dims.layers {
let p = format!("encoder.layer.{i}");
let attention_probs_dropout =
site_dropout(root_seed, &attention_probs_site(i), DROPOUT_P)?;
let attention_masks: Arc<dyn AttentionDropoutMasks> = attention_probs_dropout.clone();
let mut attention = MultiHeadAttention::new(h, dims.heads)
.with_dropout(DROPOUT_P)
.with_attention_dropout_masks(attention_masks);
install_projection(attention.q_proj_mut(), &read, &p, "query", h)?;
install_projection(attention.k_proj_mut(), &read, &p, "key", h)?;
install_projection(attention.v_proj_mut(), &read, &p, "value", h)?;
let out_proj = attention.out_proj_mut();
out_proj.set_weight(read(
&format!("{p}.attention.output.dense.weight"),
&[h, h],
)?);
out_proj.set_bias(read(&format!("{p}.attention.output.dense.bias"), &[h])?);
let mut attention_layer_norm = LayerNorm::with_eps(&[h], eps);
attention_layer_norm.set_weight(read(
&format!("{p}.attention.output.LayerNorm.weight"),
&[h],
)?);
attention_layer_norm
.set_bias(read(&format!("{p}.attention.output.LayerNorm.bias"), &[h])?);
let im = dims.intermediate;
let mut intermediate = Linear::new(h, im);
intermediate.set_weight(read(&format!("{p}.intermediate.dense.weight"), &[im, h])?);
intermediate.set_bias(read(&format!("{p}.intermediate.dense.bias"), &[im])?);
let mut output_dense = Linear::new(im, h);
output_dense.set_weight(read(&format!("{p}.output.dense.weight"), &[h, im])?);
output_dense.set_bias(read(&format!("{p}.output.dense.bias"), &[h])?);
let mut output_layer_norm = LayerNorm::with_eps(&[h], eps);
output_layer_norm.set_weight(read(&format!("{p}.output.LayerNorm.weight"), &[h])?);
output_layer_norm.set_bias(read(&format!("{p}.output.LayerNorm.bias"), &[h])?);
layers.push(EncoderLayer {
attention,
attention_probs_dropout,
attention_output_dropout: site_dropout(
root_seed,
&attention_output_site(i),
DROPOUT_P,
)?,
attention_layer_norm,
intermediate,
output_dense,
output_dropout: site_dropout(root_seed, &ffn_output_site(i), DROPOUT_P)?,
output_layer_norm,
});
}
let mut encoder = Self {
word_embeddings: read("embeddings.word_embeddings.weight", &[dims.vocab, h])?,
position_embeddings: read(
"embeddings.position_embeddings.weight",
&[dims.max_positions, h],
)?,
token_type_embeddings: read(
"embeddings.token_type_embeddings.weight",
&[dims.type_vocab, h],
)?,
embeddings_layer_norm,
embeddings_dropout: site_dropout(root_seed, EMBEDDINGS_DROPOUT_SITE, DROPOUT_P)?,
layers,
dims,
remap,
layer_norm_eps: eps,
source_revision,
tokenizer_sha256,
root_seed,
training: true,
forward_ordinal: 0,
};
encoder.set_training(false);
Ok(encoder)
}
#[must_use]
pub fn max_seq(&self) -> usize {
MAX_SEQUENCE_LENGTH.min(self.dims.max_positions)
}
#[must_use]
pub fn root_seed(&self) -> u64 {
self.root_seed
}
#[must_use]
pub fn forward_ordinal(&self) -> u64 {
self.forward_ordinal
}
pub fn set_forward_ordinal(&mut self, forward_ordinal: u64) -> Result<(), SetFitError> {
dropout_rng::checked_forward_ordinal(forward_ordinal)?;
for site in self.dropout_modules() {
site.set_forward_ordinal(forward_ordinal)?;
}
self.forward_ordinal = forward_ordinal;
Ok(())
}
fn dropout_modules(&self) -> impl Iterator<Item = &SiteDropout> {
std::iter::once(&self.embeddings_dropout)
.chain(self.layers.iter().flat_map(|l| {
[
&l.attention_probs_dropout,
&l.attention_output_dropout,
&l.output_dropout,
]
}))
.map(Arc::as_ref)
}
#[must_use]
pub fn num_layers(&self) -> usize {
self.layers.len()
}
#[must_use]
pub fn tokenizer_sha256(&self) -> &str {
&self.tokenizer_sha256
}
#[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 source_revision(&self) -> &str {
&self.source_revision
}
#[must_use]
pub fn vocab_remap(&self) -> Option<&VocabRemap> {
self.remap.as_ref()
}
#[cfg(test)]
pub(crate) fn dropout_sites(&self) -> Vec<String> {
let mut out = Vec::new();
if self.embeddings_dropout.probability() > 0.0 {
out.push(EMBEDDINGS_DROPOUT_SITE.to_string());
}
for (i, layer) in self.layers.iter().enumerate() {
if layer.attention.dropout_p() > 0.0
&& layer.attention.has_attention_dropout_masks()
&& layer.attention_probs_dropout.probability() > 0.0
{
out.push(attention_probs_site(i));
}
if layer.attention_output_dropout.probability() > 0.0 {
out.push(attention_output_site(i));
}
if layer.output_dropout.probability() > 0.0 {
out.push(ffn_output_site(i));
}
}
out
}
#[provable_contracts_macros::contract(
"setfit-encoder-conformance-v1",
equation = "setfit_encoder_forward"
)]
pub fn forward_tokens(&self, batch: &SentenceBatch) -> Result<Tensor, SetFitError> {
contract_pre_setfit_encoder_forward!(batch.input_ids());
let (_, mut layer_outputs) = self.forward_layers(batch)?;
let result = layer_outputs.pop().ok_or(SetFitError::BatchInvalid {
reason: "encoder has no layers".to_string(),
})?;
contract_post_setfit_encoder_forward!(result.data());
Ok(result)
}
#[cfg(feature = "conformance-fixtures")]
pub fn forward_tokens_per_layer(
&self,
batch: &SentenceBatch,
) -> Result<(Tensor, Vec<Tensor>), SetFitError> {
self.forward_layers(batch)
}
pub fn encode(&self, batch: &SentenceBatch) -> Result<Tensor, SetFitError> {
let tokens = self.forward_tokens(batch)?;
let pooled = masked_mean_pool(&tokens, batch.attention_mask())?;
Ok(l2_normalize_rows(&pooled, L2_EPS)?)
}
const ENCODE_BACKEND: ExecutionBackend = ExecutionBackend {
device: "cpu",
kernel: "autograd-trueno-matmul",
};
pub fn encode_with_backend(
&self,
batch: &SentenceBatch,
) -> Result<(Tensor, ExecutionBackend), SetFitError> {
let pooled = self.encode(batch)?;
Ok((pooled, Self::ENCODE_BACKEND))
}
fn forward_layers(&self, batch: &SentenceBatch) -> Result<(Tensor, Vec<Tensor>), SetFitError> {
let ids = self.validate(batch)?;
let b = batch.batch;
let s = batch.seq;
let word = embedding_gather(&self.word_embeddings, &ids, b, s)?;
let position_ids: Vec<u32> = (0..b)
.flat_map(|_| (0..s).map(|p| u32::try_from(p).unwrap_or(u32::MAX)))
.collect();
let position = embedding_gather(&self.position_embeddings, &position_ids, b, s)?;
let token_type =
embedding_gather(&self.token_type_embeddings, &batch.token_type_ids, b, s)?;
let summed = word.add(&position).add(&token_type);
let normalized = self.embeddings_layer_norm.forward(&summed);
let embeddings_out = self.embeddings_dropout.forward(&normalized);
let attention_mask = additive_attention_mask(&batch.attention_mask, b, s)?;
let mut x = embeddings_out.clone();
let mut layer_outputs = Vec::with_capacity(self.layers.len());
for layer in &self.layers {
let (attended, _) = layer.attention.forward_self(&x, Some(&attention_mask));
let attended = layer.attention_output_dropout.forward(&attended);
x = layer.attention_layer_norm.forward(&x.add(&attended));
let intermediate = layer.intermediate.forward(&x).gelu_exact();
let ffn = layer.output_dense.forward(&intermediate);
let ffn = layer.output_dropout.forward(&ffn);
x = layer.output_layer_norm.forward(&x.add(&ffn));
layer_outputs.push(x.clone());
}
Ok((embeddings_out, layer_outputs))
}
fn validate(&self, batch: &SentenceBatch) -> Result<Vec<u32>, SetFitError> {
if batch.tokenizer_sha256 != self.tokenizer_sha256 {
return Err(SetFitError::TokenizerHashMismatch {
expected: self.tokenizer_sha256.clone(),
got: batch.tokenizer_sha256.clone(),
});
}
let b = batch.batch;
let s = batch.seq;
if b == 0 || s == 0 {
return Err(SetFitError::BatchInvalid {
reason: format!("batch {b} x seq {s}: neither dimension may be zero"),
});
}
let positions = b.checked_mul(s).ok_or_else(|| SetFitError::BatchInvalid {
reason: format!("batch {b} x seq {s} overflows usize"),
})?;
for (field, len) in [
("input_ids", batch.input_ids.len()),
("token_type_ids", batch.token_type_ids.len()),
("attention_mask", batch.attention_mask.len()),
] {
if len != positions {
return Err(SetFitError::BatchInvalid {
reason: format!(
"{field} has {len} entries but batch {b} x seq {s} needs {positions}"
),
});
}
}
let max = self.max_seq();
if s > max {
return Err(SetFitError::OversizeInput { len: s, max });
}
for (position, v) in batch.attention_mask.iter().enumerate() {
if *v > 1 {
return Err(OpError::NonBinaryMaskValue {
value: *v,
position,
}
.into());
}
}
for row in 0..b {
let base = row * s;
if !batch.attention_mask[base..base + s].iter().any(|v| *v == 1) {
return Err(OpError::AllPaddingRow { row }.into());
}
}
for (position, t) in batch.token_type_ids.iter().enumerate() {
if *t as usize >= self.dims.type_vocab {
return Err(OpError::OutOfVocabulary {
id: *t,
vocab_size: self.dims.type_vocab,
position,
}
.into());
}
}
let mut rows = Vec::with_capacity(positions);
match &self.remap {
Some(remap) => {
for id in &batch.input_ids {
rows.push(remap.to_slice_row(*id)?);
}
}
None => {
for (position, id) in batch.input_ids.iter().enumerate() {
if *id as usize >= self.dims.vocab {
return Err(OpError::OutOfVocabulary {
id: *id,
vocab_size: self.dims.vocab,
position,
}
.into());
}
rows.push(*id);
}
}
}
Ok(rows)
}
fn layer_prefix(i: usize) -> String {
format!("encoder.layer.{i}")
}
}
fn install_projection<F>(
proj: &mut Linear,
read: &F,
layer_prefix: &str,
hf_name: &str,
hidden: usize,
) -> Result<(), SetFitError>
where
F: Fn(&str, &[usize]) -> Result<Tensor, SetFitError>,
{
proj.set_weight(read(
&format!("{layer_prefix}.attention.self.{hf_name}.weight"),
&[hidden, hidden],
)?);
proj.set_bias(read(
&format!("{layer_prefix}.attention.self.{hf_name}.bias"),
&[hidden],
)?);
Ok(())
}
fn hf_attention_name(local: &str) -> String {
match local.split_once('.') {
Some(("q_proj", leaf)) => format!("attention.self.query.{leaf}"),
Some(("k_proj", leaf)) => format!("attention.self.key.{leaf}"),
Some(("v_proj", leaf)) => format!("attention.self.value.{leaf}"),
Some(("out_proj", leaf)) => format!("attention.output.dense.{leaf}"),
_ => format!("attention.{local}"),
}
}
impl Module for BertSentenceEncoder {
fn forward(&self, input: &Tensor) -> Tensor {
input.clone()
}
fn parameters(&self) -> Vec<&Tensor> {
self.named_parameters()
.into_iter()
.map(|(_, t)| t)
.collect()
}
fn parameters_mut(&mut self) -> Vec<&mut Tensor> {
self.named_parameters_mut()
.into_iter()
.map(|(_, t)| t)
.collect()
}
fn named_parameters(&self) -> Vec<(String, &Tensor)> {
let mut out: Vec<(String, &Tensor)> = vec![
(
"embeddings.word_embeddings.weight".to_string(),
&self.word_embeddings,
),
(
"embeddings.position_embeddings.weight".to_string(),
&self.position_embeddings,
),
(
"embeddings.token_type_embeddings.weight".to_string(),
&self.token_type_embeddings,
),
];
out.extend(
self.embeddings_layer_norm
.named_parameters()
.into_iter()
.map(|(n, t)| (format!("embeddings.LayerNorm.{n}"), t)),
);
for (i, layer) in self.layers.iter().enumerate() {
let p = Self::layer_prefix(i);
out.extend(
layer
.attention
.named_parameters()
.into_iter()
.map(|(n, t)| (format!("{p}.{}", hf_attention_name(&n)), t)),
);
out.extend(
layer
.attention_layer_norm
.named_parameters()
.into_iter()
.map(|(n, t)| (format!("{p}.attention.output.LayerNorm.{n}"), t)),
);
out.extend(
layer
.intermediate
.named_parameters()
.into_iter()
.map(|(n, t)| (format!("{p}.intermediate.dense.{n}"), t)),
);
out.extend(
layer
.output_dense
.named_parameters()
.into_iter()
.map(|(n, t)| (format!("{p}.output.dense.{n}"), t)),
);
out.extend(
layer
.output_layer_norm
.named_parameters()
.into_iter()
.map(|(n, t)| (format!("{p}.output.LayerNorm.{n}"), t)),
);
}
out
}
fn named_parameters_mut(&mut self) -> Vec<(String, &mut Tensor)> {
let mut out: Vec<(String, &mut Tensor)> = vec![
(
"embeddings.word_embeddings.weight".to_string(),
&mut self.word_embeddings,
),
(
"embeddings.position_embeddings.weight".to_string(),
&mut self.position_embeddings,
),
(
"embeddings.token_type_embeddings.weight".to_string(),
&mut self.token_type_embeddings,
),
];
out.extend(
self.embeddings_layer_norm
.named_parameters_mut()
.into_iter()
.map(|(n, t)| (format!("embeddings.LayerNorm.{n}"), t)),
);
for (i, layer) in self.layers.iter_mut().enumerate() {
let p = Self::layer_prefix(i);
out.extend(
layer
.attention
.named_parameters_mut()
.into_iter()
.map(|(n, t)| (format!("{p}.{}", hf_attention_name(&n)), t)),
);
out.extend(
layer
.attention_layer_norm
.named_parameters_mut()
.into_iter()
.map(|(n, t)| (format!("{p}.attention.output.LayerNorm.{n}"), t)),
);
out.extend(
layer
.intermediate
.named_parameters_mut()
.into_iter()
.map(|(n, t)| (format!("{p}.intermediate.dense.{n}"), t)),
);
out.extend(
layer
.output_dense
.named_parameters_mut()
.into_iter()
.map(|(n, t)| (format!("{p}.output.dense.{n}"), t)),
);
out.extend(
layer
.output_layer_norm
.named_parameters_mut()
.into_iter()
.map(|(n, t)| (format!("{p}.output.LayerNorm.{n}"), t)),
);
}
out
}
fn set_training(&mut self, training: bool) {
self.training = training;
for site in self.dropout_modules() {
site.set_training(training);
}
self.embeddings_layer_norm.set_training(training);
for layer in &mut self.layers {
layer.attention.set_training(training);
layer.attention_layer_norm.set_training(training);
layer.intermediate.set_training(training);
layer.output_dense.set_training(training);
layer.output_layer_norm.set_training(training);
}
}
fn train(&mut self) {
self.set_training(true);
}
fn eval(&mut self) {
self.set_training(false);
}
fn training(&self) -> bool {
self.training
}
}
#[cfg(all(test, feature = "setfit"))]
#[path = "encoder_tests.rs"]
mod encoder_tests;