use alloc::format;
use ruda_model::{module::{Module,ModuleVisitor,Param},record::{PrecisionSettings,Record,Recorder,RecorderError},
tensor::{Bool,Tensor,backend::Backend}};
use crate::{LoRALinear,LoRAAdapterRecord,Dropout,pool::{pool_sequence,pool_packed_sequences,SequencePooling,SequencePoolOutput},
attention::PackedSequenceLayout};
use super::{TransformerHead,TransformerAdapterConfig,DenseTransformerNorm,SequenceHeadOutput};
#[derive(Module,Debug)]
pub struct AdaptedTransformerHead<B: Backend> {
pub projection: LoRALinear<B>,
pub normalization: Option<DenseTransformerNorm<B>>,
pub dropout: Dropout,
}
impl<B: Backend> AdaptedTransformerHead<B> {
pub fn from_dense(base: TransformerHead<B>,config: &TransformerAdapterConfig) -> Self {
let dtype = config.adapter_dtype.unwrap_or_else(||base.projection.weight.val().dtype());
Self {projection:config.lora.init_with_options(base.projection,dtype,config.use_rslora),
normalization:base.normalization,dropout:base.dropout}
}
pub fn forward<const D: usize>(&self,hidden: Tensor<B,D>) -> Tensor<B,D> {
let hidden = if let Some(norm) = &self.normalization { norm.forward(hidden) } else { hidden };
self.projection.forward(self.dropout.forward(hidden))
}
pub fn forward_sequence(&self,hidden: Tensor<B,3>,visible: Tensor<B,2,Bool>,pooling: SequencePooling)
-> SequenceHeadOutput<B> {
self.forward_pooled(pool_sequence(hidden,visible,pooling))
}
pub fn forward_packed_sequences(&self,hidden: Tensor<B,2>,layout: &PackedSequenceLayout,
visible: Option<Tensor<B,1,Bool>>,pooling: SequencePooling) -> SequenceHeadOutput<B> {
self.forward_pooled(pool_packed_sequences(hidden,layout,visible,pooling))
}
pub fn forward_pooled(&self,pooled: SequencePoolOutput<B>) -> SequenceHeadOutput<B> {
SequenceHeadOutput {logits:self.forward(pooled.values),valid_rows:pooled.valid_rows,token_counts:pooled.token_counts}
}
pub fn adapter_record(&self,base_id: &str) -> Result<TransformerHeadAdapterRecord<B>,RecorderError> {
TransformerHeadAdapterRecord::capture(self,base_id)
}
}
pub struct TransformerHeadAdapterRecord<B: Backend> {
projection: LoRAAdapterRecord<B>,
}
impl<B: Backend> Record<B> for TransformerHeadAdapterRecord<B> {
type Item<S: PrecisionSettings> = <LoRAAdapterRecord<B> as Record<B>>::Item<S>;
fn into_item<S: PrecisionSettings>(self) -> Self::Item<S> { self.projection.into_item::<S>() }
fn from_item<S: PrecisionSettings>(item: Self::Item<S>,device: &B::Device) -> Self {
Self {projection:LoRAAdapterRecord::<B>::from_item::<S>(item,device)}
}
}
fn invalid(reason: &str) -> RecorderError {
RecorderError::Unknown(format!("Invalid native head adapter record: {reason}"))
}
fn check_adapter_only<B: Backend>(head: &AdaptedTransformerHead<B>) -> Result<(),RecorderError> {
struct FrozenCheck { trainable: bool }
impl<B: Backend> ModuleVisitor<B> for FrozenCheck {
fn visit_float<const D: usize>(&mut self,param: &Param<Tensor<B,D>>) {
self.trainable |= param.val().is_require_grad();
}
}
if let Some(norm) = &head.normalization {
let mut visitor = FrozenCheck {trainable:false};
norm.visit(&mut visitor);
if visitor.trainable { return Err(invalid("trainable head normalization requires a full model checkpoint")); }
}
Ok(())
}
impl<B: Backend> TransformerHeadAdapterRecord<B> {
pub fn capture(head: &AdaptedTransformerHead<B>,base_id: &str) -> Result<Self,RecorderError> {
check_adapter_only(head)?;
Ok(Self {projection:head.projection.adapter_record(base_id)?})
}
pub fn validate_for(&self,head: &AdaptedTransformerHead<B>,base_id: &str) -> Result<(),RecorderError> {
check_adapter_only(head)?;
self.projection.schema.validate_for(&head.projection,base_id)
}
pub fn save<R: Recorder<B>>(self,recorder: &R,args: R::RecordArgs) -> Result<R::RecordOutput,RecorderError> {
recorder.record(self,args)
}
pub fn load<R: Recorder<B>>(recorder: &R,args: R::LoadArgs,device: &B::Device) -> Result<Self,RecorderError> {
recorder.load(args,device)
}
pub fn restore_into(self,mut head: AdaptedTransformerHead<B>,base_id: &str)
-> Result<AdaptedTransformerHead<B>,RecorderError> {
self.validate_for(&head,base_id)?;
head.projection = self.projection.restore_into(head.projection,base_id)?;
Ok(head)
}
}