33 macros:: * ,
44} ;
55use crate :: {
6- TPAction ,
7- weight_types:: { AttnQKV , RowTPWeight } ,
6+ Arg , TPAction ,
7+ weight_types:: { AttnQKV , ColumnTPWeight , RowTPWeight } ,
88} ;
99use tensor:: digit_layout:: types;
1010
1111#[ derive( Clone ) ]
12- pub struct Attention < T > {
12+ pub struct Attention < T : Clone > {
1313 pub nh : usize ,
1414 pub nkvh : usize ,
15- pub qkv : Linear < T > ,
15+ pub qkv : QKVFormat < T > ,
1616 pub q_norm : Option < Normalization < T > > ,
1717 pub k_norm : Option < Normalization < T > > ,
1818 pub rope : Option < RoPE < T > > ,
1919 pub output : Linear < T > ,
2020}
2121
22+ #[ derive( Clone ) ]
23+ pub enum QKVFormat < T : Clone > {
24+ Combined ( Linear < T > ) ,
25+ Separated {
26+ q : Linear < T > ,
27+ k : Linear < T > ,
28+ v : Linear < T > ,
29+ } ,
30+ }
31+
2232#[ derive( Clone ) ]
2333pub struct RoPE < T > {
2434 pub multimodal : bool ,
@@ -27,7 +37,7 @@ pub struct RoPE<T> {
2737 pub cos : T ,
2838}
2939
30- impl < T > Attention < T > {
40+ impl < T : Clone > Attention < T > {
3141 pub fn tensor_parallel ( self , dist : Distribution ) -> Attention < TPTensor < T > > {
3242 let Self {
3343 nh,
@@ -43,7 +53,16 @@ impl<T> Attention<T> {
4353 Attention {
4454 nh : nh / dist. total * dist. len ,
4555 nkvh : nkvh / dist. total * dist. len ,
46- qkv : qkv. parallel ( TPAction :: new ( AttnQKV ( nh / nkvh) , dist) ) ,
56+ qkv : match qkv {
57+ QKVFormat :: Combined ( qkv) => {
58+ QKVFormat :: Combined ( qkv. parallel ( TPAction :: new ( AttnQKV ( nh / nkvh) , dist) ) )
59+ }
60+ QKVFormat :: Separated { q, k, v } => QKVFormat :: Separated {
61+ q : q. parallel ( TPAction :: new ( ColumnTPWeight , dist) ) ,
62+ k : k. parallel ( TPAction :: new ( ColumnTPWeight , dist) ) ,
63+ v : v. parallel ( TPAction :: new ( ColumnTPWeight , dist) ) ,
64+ } ,
65+ } ,
4766 q_norm : q_norm. map ( |norm| norm. tensor_parallel ( ) ) ,
4867 k_norm : k_norm. map ( |norm| norm. tensor_parallel ( ) ) ,
4968 rope : rope. map (
@@ -64,7 +83,7 @@ impl<T> Attention<T> {
6483 }
6584}
6685
67- impl < T > NuralNetwork < T > for Attention < T > {
86+ impl < T : Clone > NuralNetwork < T > for Attention < T > {
6887 fn launch (
6988 self ,
7089 inputs : impl IntoIterator < Item = Tensor < T > > ,
@@ -81,12 +100,36 @@ impl<T> NuralNetwork<T> for Attention<T> {
81100 rope,
82101 output,
83102 } = self ;
84- destruct ! ( [ x] = ctx. trap( "attn-qkv" , qkv, [ x] ) ?) ;
85- dims ! ( [ _, dqkv] = x) ;
86- let dh = dqkv. clone ( ) / ( nh + nkvh + nkvh) ;
87103
88- destruct ! ( [ q, k, v] = x. split( "split-qkv" , 1 , [ nh. into( ) , nkvh. into( ) , nkvh. into( ) ] ) ?) ;
104+ dims ! ( [ _, d] = x) ;
105+ let dh = d. clone ( ) / nh;
89106
107+ let [ q, k, v] = match qkv {
108+ QKVFormat :: Combined ( qkv) => {
109+ destruct ! ( [ x] = ctx. trap( "attn-qkv" , qkv, [ x] ) ?) ;
110+ destruct ! (
111+ [ q, k, v] = ctx. call(
112+ "split-qkv" ,
113+ "split" ,
114+ Some ( Arg :: dict( [
115+ ( "axis" . into( ) , Arg :: int( 1 ) ) ,
116+ (
117+ "parts" . into( ) ,
118+ Arg :: arr( [ Arg :: dim( nh) , Arg :: dim( nkvh) , Arg :: dim( nkvh) ] )
119+ )
120+ ] ) ) ,
121+ [ x] ,
122+ ) ?
123+ ) ;
124+ [ q, k, v]
125+ }
126+ QKVFormat :: Separated { q, k, v } => {
127+ destruct ! ( [ q] = ctx. trap( "attn-q" , q, [ x. clone( ) ] ) ?) ;
128+ destruct ! ( [ k] = ctx. trap( "attn-k" , k, [ x. clone( ) ] ) ?) ;
129+ destruct ! ( [ v] = ctx. trap( "attn-v" , v, [ x] ) ?) ;
130+ [ q, k, v]
131+ }
132+ } ;
90133 // Apply normalization to q and k if they exist
91134 let q = match q_norm {
92135 Some ( norm) => {
@@ -114,9 +157,8 @@ impl<T> NuralNetwork<T> for Attention<T> {
114157 cos,
115158 } ) => {
116159 let shape = [ nctx. into ( ) , dh. clone ( ) / 2 ] ;
117- let sin = ctx. load_external ( "rope.sin" , types:: F32 , shape. clone ( ) , sin) ;
118- let cos = ctx. load_external ( "rope.cos" , types:: F32 , shape, cos) ;
119-
160+ destruct ! ( [ sin] = ctx. load_external( "rope.sin" , types:: F32 , shape. clone( ) , sin) ?) ;
161+ destruct ! ( [ cos] = ctx. load_external( "rope.cos" , types:: F32 , shape, cos) ?) ;
120162 let op = if multimodal { "mrope" } else { "rope" } ;
121163 destruct ! (
122164 [ q_] = ctx. call(
0 commit comments