use super::{Backend,Tensor,Param,check_storage,ColumnParallelLinear};
use crate::Embedding;
use super::super::{VocabParallelLossLayout,VocabParallelEmbedding,VocabParallelProjection,VocabParallelTransformerHead,
VocabParallelLoRAProjection,VocabParallelAdaptedTransformerHead};
fn rows<B:Backend,const D:usize>(parameter:Param<Tensor<B,D>>,layout:&VocabParallelLossLayout,rank:usize) -> Param<Tensor<B,D>> {
assert!(D > 0,"vocabulary storage requires a row axis");
let shape = parameter.val().dims();let interval = layout.interval(rank);
assert_eq!(shape[0],layout.storage_size(),"full vocabulary rows differ from declared actual storage");
if interval == (0..shape[0]) {return parameter;}
check_storage(¶meter.val());
parameter.map(|value| {
let trainable = value.is_require_grad();value.slice_dim(0,interval).detach().set_require_grad(trainable)
})
}
fn same_embedding<B:Backend>(embedding:&Embedding<B>,projection:&VocabParallelProjection<B>) {
let source = embedding.weight.val();let target = projection.weight.val();
assert!(embedding.weight.id == projection.weight.id && source.dims() == target.dims() && source.dtype() == target.dtype()
&& source.device() == target.device() && source.is_require_grad() == target.is_require_grad(),
"full embedding/head must identify the same actual native parameter and contract");
}
impl<B:Backend> VocabParallelEmbedding<B> {
pub fn from_full_with_layout(mut embedding:Embedding<B>,layout:&VocabParallelLossLayout,rank:usize,padding_index:Option<usize>) -> Self {
assert!(embedding.weight.val().dims()[1] > 0,"full embedding hidden width must be positive");
embedding.weight = rows(embedding.weight,layout,rank);
Self::from_shard_with_layout(embedding,layout,rank,padding_index)
}
}
impl<B:Backend> VocabParallelProjection<B> {
pub fn from_full_with_layout(self,layout:&VocabParallelLossLayout,rank:usize) -> Self {
assert!(self.weight.val().dims()[1] > 0,"full vocabulary head hidden width must be positive");
if let Some(bias) = &self.bias {
assert_eq!(bias.val().dims(),[layout.storage_size()],"full vocabulary head bias rows differ");
assert_eq!(bias.val().device(),self.weight.val().device(),"full vocabulary head bias/weight devices differ");
}
Self {weight:rows(self.weight,layout,rank),bias:self.bias.map(|bias|rows(bias,layout,rank))}
}
}
impl<B:Backend> VocabParallelTransformerHead<B> {
pub fn from_full(head:Self,layout:&VocabParallelLossLayout,rank:usize) -> Self {
Self::from_projection(VocabParallelProjection::from_full_with_layout(head.projection,layout,rank),head.normalization,head.dropout,layout,rank)
}
pub fn from_full_tied(embedding:Embedding<B>,head:Self,layout:&VocabParallelLossLayout,rank:usize,padding_index:Option<usize>)
-> (VocabParallelEmbedding<B>,Self) {
same_embedding(&embedding,&head.projection);
let embedding = VocabParallelEmbedding::from_full_with_layout(embedding,layout,rank,padding_index);
let bias = head.projection.bias.map(|bias|rows(bias,layout,rank));
let weight = head.projection.weight.map(|_|embedding.local.weight.val());
let head = Self::from_projection(VocabParallelProjection {weight,bias},head.normalization,head.dropout,layout,rank);
(embedding,head)
}
}
impl<B:Backend> VocabParallelLoRAProjection<B> {
pub fn from_full(mut layer:Self,layout:&VocabParallelLossLayout,rank:usize) -> Self {
let [classes,hidden] = layer.base.weight.val().dims();let adapter_rank = layer.adapter_a.weight.val().dims()[1];
assert_eq!(classes,layout.storage_size(),"full vocabulary adapter base rows differ");
assert_eq!(layer.adapter_a.weight.val().dims(),[hidden,adapter_rank],"full vocabulary adapter A geometry differs");
assert_eq!(layer.adapter_b.weight.val().dims(),[adapter_rank,classes],"full vocabulary adapter B geometry differs");
let interval = layout.interval(rank);
if interval != (0..classes) {
assert!(layer.base.weight.id != layer.adapter_a.weight.id && layer.base.weight.id != layer.adapter_b.weight.id
&& layer.adapter_a.weight.id != layer.adapter_b.weight.id,"tied full adapter roles require explicit compatible local loading");
}
layer.base = VocabParallelProjection::from_full_with_layout(layer.base,layout,rank);
layer.adapter_b = ColumnParallelLinear::from_full(layer.adapter_b,interval).local;
Self::from_adapters(layer.base,layer.adapter_a,layer.adapter_b,layer.dropout,layer.scale)
}
}
impl<B:Backend> VocabParallelAdaptedTransformerHead<B> {
pub fn from_full(head:Self,layout:&VocabParallelLossLayout,rank:usize) -> Self {
Self::from_projection(VocabParallelLoRAProjection::from_full(head.projection,layout,rank),head.normalization,head.dropout,layout,rank)
}
pub fn from_full_tied(embedding:Embedding<B>,head:Self,layout:&VocabParallelLossLayout,rank:usize,padding_index:Option<usize>)
-> (VocabParallelEmbedding<B>,Self) {
same_embedding(&embedding,&head.projection.base);
let full = head.projection;
let classes = full.base.weight.val().dims()[0];let hidden = full.base.weight.val().dims()[1];let adapter_rank = full.adapter_a.weight.val().dims()[1];
assert_eq!(full.adapter_a.weight.val().dims(),[hidden,adapter_rank],"full tied head adapter A geometry differs");
assert_eq!(full.adapter_b.weight.val().dims(),[adapter_rank,classes],"full tied head adapter B geometry differs");
let interval = layout.interval(rank);
if interval != (0..classes) {
assert!(full.base.weight.id != full.adapter_a.weight.id && full.base.weight.id != full.adapter_b.weight.id
&& full.adapter_a.weight.id != full.adapter_b.weight.id,"tied full adapter roles require compatible explicit local loading");
}
let embedding = VocabParallelEmbedding::from_full_with_layout(embedding,layout,rank,padding_index);
let base = VocabParallelProjection {weight:full.base.weight.map(|_|embedding.local.weight.val()),bias:full.base.bias.map(|bias|rows(bias,layout,rank))};
let projection = VocabParallelLoRAProjection::from_adapters(base,full.adapter_a,
ColumnParallelLinear::from_full(full.adapter_b,interval).local,full.dropout,full.scale);
let head = Self::from_projection(projection,head.normalization,head.dropout,layout,rank);
(embedding,head)
}
}