///|
/// A BERT encoder layer using the original post-normalization order.
///
/// Dropout is intentionally absent because this model represents inference.
/// SelfAttention and FeedForward remain reusable pure sublayers; this block
/// owns the residual and LayerNorm boundaries required by BERT checkpoints.
pub(all) struct BertEncoderBlock[T] {
attention : SelfAttention[T]
attention_output_normalization : LayerNorm[T]
feed_forward : FeedForward[T]
output_normalization : LayerNorm[T]
}
///|
pub fn[T : @tensor.TensorOps] BertEncoderBlock::forward(
self : BertEncoderBlock[T],
input : T,
) -> T raise {
let input_shape = input.shape()
let attention_output = self.attention.forward(input)
require_residual_shape("BERT attention", attention_output, input_shape)
let attention_residual = self.attention_output_normalization.forward(
input.add(attention_output),
)
let feed_forward_output = self.feed_forward.forward(attention_residual)
require_residual_shape("BERT feed-forward", feed_forward_output, input_shape)
self.output_normalization.forward(attention_residual.add(feed_forward_output))
}
///|
/// A non-empty sequence of post-normalization BERT encoder layers.
pub struct BertEncoderStack[T] {
layers_ : Array[BertEncoderBlock[T]]
}
///|
pub fn[T] BertEncoderStack::new(
layers : Array[BertEncoderBlock[T]],
) -> BertEncoderStack[T] raise @tensor.TensorError {
if layers.length() == 0 {
raise @tensor.TensorError::new(
"BERT encoder stack must contain at least one layer",
)
}
{ layers_: layers.copy() }
}
///|
pub fn[T] BertEncoderStack::layer_count(self : BertEncoderStack[T]) -> Int {
self.layers_.length()
}
///|
pub fn[T : @tensor.TensorOps] BertEncoderStack::forward(
self : BertEncoderStack[T],
input : T,
) -> T raise {
let mut output = input
for layer in self.layers_ {
output = layer.forward(output)
}
output
}
///|
/// Shape and numerical configuration for an inference-only BERT encoder.
pub struct BertEncoderConfig {
encoder_ : TransformerEncoderConfig
layers_ : Int
}
///|
pub fn BertEncoderConfig::new(
layers : Int,
width : Int,
heads : Int,
intermediate_size : Int,
input_rank : Int,
epsilon : Float,
) -> BertEncoderConfig raise @tensor.TensorError {
if layers <= 0 {
raise @tensor.TensorError::new("BERT encoder layers must be positive")
}
{
encoder_: TransformerEncoderConfig::new(
width, heads, intermediate_size, input_rank, epsilon,
),
layers_: layers,
}
}
///|
pub fn BertEncoderConfig::encoder(
self : BertEncoderConfig,
) -> TransformerEncoderConfig {
self.encoder_
}
///|
pub fn BertEncoderConfig::layers(self : BertEncoderConfig) -> Int {
self.layers_
}
///|
pub fn BertEncoderConfig::width(self : BertEncoderConfig) -> Int {
self.encoder_.width()
}
///|
pub fn BertEncoderConfig::heads(self : BertEncoderConfig) -> Int {
self.encoder_.heads()
}
///|
pub fn BertEncoderConfig::intermediate_size(self : BertEncoderConfig) -> Int {
self.encoder_.hidden_size()
}
///|
/// Owned per-layer parameters for a BERT encoder stack.
pub struct BertEncoderParameters {
layers_ : Array[TransformerEncoderParameters]
}
///|
pub fn BertEncoderParameters::new(
config : BertEncoderConfig,
layers : Array[TransformerEncoderParameters],
) -> BertEncoderParameters raise @tensor.TensorError {
if layers.length() != config.layers() {
raise @tensor.TensorError::new(
"BERT encoder parameter count " +
layers.length().to_string() +
" must match " +
config.layers().to_string() +
" layers",
)
}
let encoder_config = config.encoder()
for index, layer in layers {
if !layer.matches(encoder_config) {
raise @tensor.TensorError::new(
"BERT encoder layer " +
index.to_string() +
" parameters do not match config",
)
}
}
{ layers_: layers.copy() }
}
///|
pub fn BertEncoderParameters::layers(
self : BertEncoderParameters,
) -> Array[TransformerEncoderParameters] {
self.layers_.copy()
}