use ruda_model::tensor::{Bool, Int, Tensor, backend::Backend};
use crate::{attention::{PackedSequenceLayout, CompressedAttentionProjection}, loss::{CausalLanguageModel, PackedCausalLanguageModel, CausalCrossEntropyConfig, CausalLoss}};
use super::HybridAttentionLanguageModel;
#[derive(Debug)]
pub struct HybridLanguageTrainingOutput<B: Backend> {
pub causal: CausalLoss<B>,
pub indexer_loss: Tensor<B, 1>,
}
#[derive(Debug)]
pub struct HybridPackedLanguageTrainingOutput<B: Backend> {
pub causal: CausalLoss<B>,
pub document_indexer_losses: Tensor<B, 1>,
}
impl<B: Backend, P: CompressedAttentionProjection<B>> CausalLanguageModel<B> for HybridAttentionLanguageModel<B, P> {
fn forward_hidden(&self, tokens: Tensor<B, 2, Int>) -> Tensor<B, 3> {
self.backbone.forward(tokens, None)
}
fn project(&self, hidden: Tensor<B, 2>) -> Tensor<B, 2> { self.project_tokens(hidden) }
}
impl<B: Backend, P: CompressedAttentionProjection<B>> PackedCausalLanguageModel<B> for HybridAttentionLanguageModel<B, P> {
fn forward_packed_hidden(&self, tokens: Tensor<B, 1, Int>, layout: &PackedSequenceLayout) -> Tensor<B, 2> {
self.backbone.forward_packed(tokens, layout, None)
}
fn project(&self, hidden: Tensor<B, 2>) -> Tensor<B, 2> { self.project_tokens(hidden) }
}
impl<B: Backend, P: CompressedAttentionProjection<B>> HybridAttentionLanguageModel<B, P> {
pub fn causal_loss(&self, tokens: Tensor<B, 2, Int>, labels: Tensor<B, 2, Int>,
valid: Option<Tensor<B, 2, Bool>>, loss: &CausalCrossEntropyConfig) -> CausalLoss<B> {
assert_eq!(tokens.dims(), labels.dims(), "hybrid training token/label geometry differs");
assert_eq!(tokens.device(), labels.device(), "hybrid training label device differs");
loss.forward_hidden(self.backbone.forward(tokens, valid), labels, |rows| self.project_tokens(rows))
}
pub fn training_forward(&self, tokens: Tensor<B, 2, Int>, labels: Tensor<B, 2, Int>,
valid: Option<Tensor<B, 2, Bool>>, loss: &CausalCrossEntropyConfig,
indexer_warmup: bool) -> HybridLanguageTrainingOutput<B> {
assert_eq!(tokens.dims(), labels.dims(), "hybrid training token/label geometry differs");
assert_eq!(tokens.device(), labels.device(), "hybrid training label device differs");
let result = self.backbone.forward_with_aux(tokens, valid, indexer_warmup);
HybridLanguageTrainingOutput {
causal: loss.forward_hidden(result.output, labels, |rows| self.project_tokens(rows)),
indexer_loss: result.indexer_loss,
}
}
pub fn packed_causal_loss(&self, tokens: Tensor<B, 1, Int>, labels: Tensor<B, 1, Int>, layout: &PackedSequenceLayout,
valid: Option<Tensor<B, 1, Bool>>, loss: &CausalCrossEntropyConfig) -> CausalLoss<B> {
assert_eq!(tokens.dims(), labels.dims(), "packed hybrid training token/label geometry differs");
assert_eq!(tokens.device(), labels.device(), "packed hybrid training label device differs");
loss.forward_packed_hidden(self.backbone.forward_packed(tokens, layout, valid), labels, layout, |rows| self.project_tokens(rows))
}
pub fn packed_training_forward(&self, tokens: Tensor<B, 1, Int>, labels: Tensor<B, 1, Int>, layout: &PackedSequenceLayout,
valid: Option<Tensor<B, 1, Bool>>, loss: &CausalCrossEntropyConfig,
indexer_warmup: bool) -> HybridPackedLanguageTrainingOutput<B> {
assert_eq!(tokens.dims(), labels.dims(), "packed hybrid training token/label geometry differs");
assert_eq!(tokens.device(), labels.device(), "packed hybrid training label device differs");
let result = self.backbone.forward_packed_with_aux(tokens, layout, valid, indexer_warmup);
HybridPackedLanguageTrainingOutput {
causal: loss.forward_packed_hidden(result.output, labels, layout, |rows| self.project_tokens(rows)),
document_indexer_losses: result.document_indexer_losses,
}
}
}