use super::{AdaptedFeedForward, AdaptedTransformerBlock, AdaptedStackLayer, AdaptedTransformerStack,
AwqFeedForward, AwqTransformerBlock, AwqTransformerStack,
DenseFeedForward, DenseTransformerBlock, DenseTransformerStack, TransformerProjection};
use super::dense::try_residual_branch;
use core::fmt;
use ruda_model::tensor::{NativeSwiGluOps, Tensor};
use crate::attention::{DenseAttentionMask, DenseAttentionOptions, PackedSequenceLayout,
PackedAttentionOptions, PackedDocumentAttentionMask};
#[derive(Debug)]
pub enum NativeFeedForwardError<E, A> {
Execution(E),
Activation(A),
}
impl<E: fmt::Debug, A: fmt::Debug> fmt::Display for NativeFeedForwardError<E, A> {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match self {
Self::Execution(error) => write!(f, "native feed-forward execution: {error:?}"),
Self::Activation(error) => write!(f, "native feed-forward activation: {error:?}"),
}
}
}
impl<E: fmt::Debug, A: fmt::Debug> core::error::Error for NativeFeedForwardError<E, A> {}
impl<B: NativeSwiGluOps> DenseFeedForward<B> {
pub fn try_forward_native<const D: usize>(&self, input: Tensor<B, D>) -> Result<Tensor<B, D>, B::SwiGluError> {
let up = self.up.forward(input.clone());
let value = if let Some(gate) = &self.gate {
self.activation.try_forward_gated_native(gate.forward(input), up)?
} else { self.activation.try_forward_native(up)? };
Ok(self.down.forward(self.dropout.forward(value)))
}
}
impl<B: NativeSwiGluOps> AdaptedFeedForward<B> {
pub fn try_forward_native<const D: usize>(&self, input: Tensor<B, D>) -> Result<Tensor<B, D>, B::SwiGluError> {
let up = self.up.forward(input.clone());
let value = if let Some(gate) = &self.gate {
self.activation.try_forward_gated_native(gate.forward(input), up)?
} else { self.activation.try_forward_native(up)? };
Ok(self.down.forward(self.dropout.forward(value)))
}
}
impl<B: NativeSwiGluOps, P: TransformerProjection<B>> AwqFeedForward<B, P> {
pub fn try_forward_native<const D: usize>(&self, input: Tensor<B, D>)
-> Result<Tensor<B, D>, NativeFeedForwardError<P::Error, B::SwiGluError>> {
let up = self.up.forward(input.clone()).map_err(NativeFeedForwardError::Execution)?;
let value = if let Some(gate) = &self.gate {
let gate = gate.forward(input).map_err(NativeFeedForwardError::Execution)?;
self.activation.try_forward_gated_native(gate, up)
} else { self.activation.try_forward_native(up) }.map_err(NativeFeedForwardError::Activation)?;
self.down.forward(self.dropout.forward(value)).map_err(NativeFeedForwardError::Execution)
}
}
impl<B: NativeSwiGluOps> DenseTransformerBlock<B> {
pub fn try_forward_feed_forward_native<const D: usize>(&self, hidden: Tensor<B, D>)
-> Result<Tensor<B, D>, B::SwiGluError> {
try_residual_branch(hidden, &self.feed_forward_norm, &self.residual_dropout, self.norm_first,
|source| self.feed_forward.try_forward_native(source))
}
pub fn try_forward_native_with_positions<F>(&self, input: Tensor<B, 3>, masks: DenseAttentionMask<B>,
options: DenseAttentionOptions, positions: F) -> Result<Tensor<B, 3>, B::SwiGluError>
where F: FnOnce(Tensor<B, 4>, Tensor<B, 4>) -> (Tensor<B, 4>, Tensor<B, 4>) {
let hidden = self.forward_attention_with_positions(input, masks, options, positions);
self.try_forward_feed_forward_native(hidden)
}
}
impl<B: NativeSwiGluOps> AdaptedTransformerBlock<B> {
pub fn try_forward_feed_forward_native<const D: usize>(&self, hidden: Tensor<B, D>) -> Result<Tensor<B, D>, B::SwiGluError> {
try_residual_branch(hidden, &self.feed_forward_norm, &self.residual_dropout, self.norm_first,
|source| self.feed_forward.try_forward_native(source))
}
pub fn try_forward_native_with_positions<F>(&self, input: Tensor<B, 3>, masks: DenseAttentionMask<B>,
options: DenseAttentionOptions, positions: F) -> Result<Tensor<B, 3>, B::SwiGluError>
where F: FnOnce(Tensor<B, 4>, Tensor<B, 4>) -> (Tensor<B, 4>, Tensor<B, 4>) {
let hidden = self.forward_attention_with_positions(input, masks, options, positions);
self.try_forward_feed_forward_native(hidden)
}
}
impl<B: NativeSwiGluOps> DenseTransformerBlock<B> {
pub fn try_forward_packed_native_with_positions<F>(&self, input: Tensor<B, 2>, layout: &PackedSequenceLayout,
masks: Option<&[PackedDocumentAttentionMask<B>]>, options: PackedAttentionOptions, positions: F)
-> Result<Tensor<B, 2>, B::SwiGluError>
where F: FnOnce(Tensor<B, 3>, Tensor<B, 3>) -> (Tensor<B, 3>, Tensor<B, 3>) {
let hidden = match masks {
Some(masks) => self.forward_packed_attention_masked(input, layout, masks, options, positions),
None => self.forward_packed_attention(input, layout, options, positions),
};
self.try_forward_feed_forward_native(hidden)
}
}
impl<B: NativeSwiGluOps> AdaptedTransformerBlock<B> {
pub fn try_forward_packed_native_with_positions<F>(&self, input: Tensor<B, 2>, layout: &PackedSequenceLayout,
masks: Option<&[PackedDocumentAttentionMask<B>]>, options: PackedAttentionOptions, positions: F)
-> Result<Tensor<B, 2>, B::SwiGluError>
where F: FnOnce(Tensor<B, 3>, Tensor<B, 3>) -> (Tensor<B, 3>, Tensor<B, 3>) {
let hidden = match masks {
Some(masks) => self.forward_packed_attention_masked(input, layout, masks, options, positions),
None => self.forward_packed_attention(input, layout, options, positions),
};
self.try_forward_feed_forward_native(hidden)
}
}
impl<B: NativeSwiGluOps, P: TransformerProjection<B>> AwqTransformerBlock<B, P> {
pub fn try_forward_packed_native_with_positions<F>(&self, input: Tensor<B, 2>, layout: &PackedSequenceLayout,
masks: Option<&[PackedDocumentAttentionMask<B>]>, options: PackedAttentionOptions, positions: F)
-> Result<Tensor<B, 2>, NativeFeedForwardError<P::Error, B::SwiGluError>>
where F: FnOnce(Tensor<B, 3>, Tensor<B, 3>) -> (Tensor<B, 3>, Tensor<B, 3>) {
let hidden = self.forward_packed_attention_with_positions(input, layout, masks, options, positions)
.map_err(NativeFeedForwardError::Execution)?;
self.try_forward_feed_forward_native(hidden)
}
}
impl<B: NativeSwiGluOps> AdaptedStackLayer<B> {
pub fn try_forward_native_with_positions<F>(&self, input: Tensor<B, 3>, masks: DenseAttentionMask<B>,
options: DenseAttentionOptions, positions: F) -> Result<Tensor<B, 3>, B::SwiGluError>
where F: FnOnce(Tensor<B, 4>, Tensor<B, 4>) -> (Tensor<B, 4>, Tensor<B, 4>) {
match self {
Self::Dense(block) => block.try_forward_native_with_positions(input, masks, options, positions),
Self::Adapted(block) => block.try_forward_native_with_positions(input, masks, options, positions),
}
}
pub fn try_forward_packed_native_with_positions<F>(&self, input: Tensor<B, 2>, layout: &PackedSequenceLayout,
masks: Option<&[PackedDocumentAttentionMask<B>]>, options: PackedAttentionOptions, positions: F)
-> Result<Tensor<B, 2>, B::SwiGluError>
where F: FnOnce(Tensor<B, 3>, Tensor<B, 3>) -> (Tensor<B, 3>, Tensor<B, 3>) {
match self {
Self::Dense(block) => block.try_forward_packed_native_with_positions(input, layout, masks, options, positions),
Self::Adapted(block) => block.try_forward_packed_native_with_positions(input, layout, masks, options, positions),
}
}
}
impl<B: NativeSwiGluOps> DenseTransformerStack<B> {
pub fn try_forward_native_with_positions<F>(&self, mut input: Tensor<B, 3>, masks: DenseAttentionMask<B>,
options: DenseAttentionOptions, mut positions: F) -> Result<Tensor<B, 3>, B::SwiGluError>
where F: FnMut(usize, Tensor<B, 4>, Tensor<B, 4>) -> (Tensor<B, 4>, Tensor<B, 4>) {
for (index, block) in self.blocks.iter().enumerate() {
input = block.try_forward_native_with_positions(input, masks.clone(), options, |query, key| positions(index, query, key))?;
}
Ok(input)
}
pub fn try_forward_packed_native_with_positions<F>(&self, mut input: Tensor<B, 2>, layout: &PackedSequenceLayout,
masks: Option<&[PackedDocumentAttentionMask<B>]>, options: PackedAttentionOptions, mut positions: F)
-> Result<Tensor<B, 2>, B::SwiGluError>
where F: FnMut(usize, Tensor<B, 3>, Tensor<B, 3>) -> (Tensor<B, 3>, Tensor<B, 3>) {
assert_eq!(input.dims()[0], layout.tokens(), "packed transformer boundaries differ from actual token rows");
for (index, block) in self.blocks.iter().enumerate() {
input = block.try_forward_packed_native_with_positions(input, layout, masks, options, |query, key| positions(index, query, key))?;
}
Ok(input)
}
}
impl<B: NativeSwiGluOps> AdaptedTransformerStack<B> {
pub fn try_forward_native_with_positions<F>(&self, mut input: Tensor<B, 3>, masks: DenseAttentionMask<B>,
options: DenseAttentionOptions, mut positions: F) -> Result<Tensor<B, 3>, B::SwiGluError>
where F: FnMut(usize, Tensor<B, 4>, Tensor<B, 4>) -> (Tensor<B, 4>, Tensor<B, 4>) {
for (index, layer) in self.layers.iter().enumerate() {
input = layer.try_forward_native_with_positions(input, masks.clone(), options, |query, key| positions(index, query, key))?;
}
Ok(input)
}
pub fn try_forward_packed_native_with_positions<F>(&self, mut input: Tensor<B, 2>, layout: &PackedSequenceLayout,
masks: Option<&[PackedDocumentAttentionMask<B>]>, options: PackedAttentionOptions, mut positions: F)
-> Result<Tensor<B, 2>, B::SwiGluError>
where F: FnMut(usize, Tensor<B, 3>, Tensor<B, 3>) -> (Tensor<B, 3>, Tensor<B, 3>) {
assert_eq!(input.dims()[0], layout.tokens(), "packed transformer boundaries differ from actual token rows");
for (index, layer) in self.layers.iter().enumerate() {
input = layer.try_forward_packed_native_with_positions(input, layout, masks, options, |query, key| positions(index, query, key))?;
}
Ok(input)
}
}
impl<B: NativeSwiGluOps, P: TransformerProjection<B>> AwqTransformerStack<B, P> {
pub fn try_forward_native_with_positions<F>(&self, mut input: Tensor<B, 3>, masks: DenseAttentionMask<B>,
options: DenseAttentionOptions, mut positions: F)
-> Result<Tensor<B, 3>, NativeFeedForwardError<P::Error, B::SwiGluError>>
where F: FnMut(usize, Tensor<B, 4>, Tensor<B, 4>) -> (Tensor<B, 4>, Tensor<B, 4>) {
for (index, block) in self.blocks.iter().enumerate() {
input = block.try_forward_native_with_positions(input, masks.clone(), options, |query, key| positions(index, query, key))?;
}
Ok(input)
}
pub fn try_forward_packed_native_with_positions<F>(&self, mut input: Tensor<B, 2>, layout: &PackedSequenceLayout,
masks: Option<&[PackedDocumentAttentionMask<B>]>, options: PackedAttentionOptions, mut positions: F)
-> Result<Tensor<B, 2>, NativeFeedForwardError<P::Error, B::SwiGluError>>
where F: FnMut(usize, Tensor<B, 3>, Tensor<B, 3>) -> (Tensor<B, 3>, Tensor<B, 3>) {
for (index, block) in self.blocks.iter().enumerate() {
input = block.try_forward_packed_native_with_positions(input, layout, masks, options, |query, key| positions(index, query, key))?;
}
Ok(input)
}
}
impl<B: NativeSwiGluOps, P: TransformerProjection<B>> AwqTransformerBlock<B, P> {
pub fn try_forward_feed_forward_native<const D: usize>(&self, hidden: Tensor<B, D>)
-> Result<Tensor<B, D>, NativeFeedForwardError<P::Error, B::SwiGluError>> {
try_residual_branch(hidden, &self.feed_forward_norm, &self.residual_dropout, self.norm_first,
|source| self.feed_forward.try_forward_native(source))
}
pub fn try_forward_native_with_positions<F>(&self, input: Tensor<B, 3>, masks: DenseAttentionMask<B>,
options: DenseAttentionOptions, positions: F)
-> Result<Tensor<B, 3>, NativeFeedForwardError<P::Error, B::SwiGluError>>
where F: FnOnce(Tensor<B, 4>, Tensor<B, 4>) -> (Tensor<B, 4>, Tensor<B, 4>) {
let hidden = self.forward_attention_with_positions(input, masks, options, positions)
.map_err(NativeFeedForwardError::Execution)?;
self.try_forward_feed_forward_native(hidden)
}
}