@@ -539,6 +539,59 @@ def get_test_model_exhaustive() -> Model:
539539 num_heads = 2 , key_dim = 3 , value_dim = 5 ,
540540 use_bias = True , output_shape = None , attention_axes = None )(inputs [49 ], inputs [50 ], inputs [51 ]))
541541
542+ # GroupedQueryAttention: shared Q/K/V seq, separate K/V seq, with/without gate.
543+ outputs .append (GroupQueryAttention (head_dim = 4 , num_query_heads = 6 ,
544+ num_key_value_heads = 2 )(inputs [49 ], inputs [49 ], inputs [49 ]))
545+ outputs .append (GroupQueryAttention (head_dim = 4 , num_query_heads = 6 ,
546+ num_key_value_heads = 2 , use_gate = True )(inputs [49 ], inputs [49 ], inputs [49 ]))
547+ outputs .append (GroupQueryAttention (head_dim = 4 , num_query_heads = 4 ,
548+ num_key_value_heads = 4 )(inputs [49 ], inputs [49 ], inputs [49 ]))
549+ outputs .append (GroupQueryAttention (head_dim = 3 , num_query_heads = 4 ,
550+ num_key_value_heads = 2 )(inputs [49 ], inputs [50 ], inputs [51 ]))
551+
552+ # DepthwiseConv1D variants on rank-2 sequence input.
553+ outputs .append (DepthwiseConv1D (kernel_size = 3 , padding = 'same' )(inputs [49 ]))
554+ outputs .append (DepthwiseConv1D (kernel_size = 2 , padding = 'valid' , strides = 2 )(inputs [49 ]))
555+ outputs .append (DepthwiseConv1D (kernel_size = 3 , padding = 'same' , use_bias = False )(inputs [49 ]))
556+
557+ # RMSNormalization / GroupNormalization on a (T, F=4) input.
558+ outputs .append (RMSNormalization ()(inputs [49 ]))
559+ outputs .append (RMSNormalization (epsilon = 1e-3 )(inputs [49 ]))
560+ outputs .append (GroupNormalization (groups = 2 )(inputs [49 ]))
561+ outputs .append (GroupNormalization (groups = 2 , scale = False )(inputs [49 ]))
562+ outputs .append (GroupNormalization (groups = 4 , center = False , epsilon = 1e-4 )(inputs [49 ]))
563+ outputs .append (GroupNormalization (groups = 1 )(inputs [49 ])) # = LayerNormalization
564+ outputs .append (GroupNormalization (groups = 4 )(inputs [49 ])) # = InstanceNormalization
565+
566+ # EinsumDense: Dense-equivalent, multi-head projection, no-bias, chained collapse.
567+ outputs .append (EinsumDense ('abc,cd->abd' , output_shape = (None , 8 ),
568+ bias_axes = 'd' )(inputs [49 ]))
569+ outputs .append (EinsumDense ('abc,cde->abde' , output_shape = (None , 3 , 4 ),
570+ bias_axes = 'de' )(inputs [49 ]))
571+ outputs .append (EinsumDense ('abc,cd->abd' , output_shape = (None , 6 ))(inputs [49 ]))
572+ outputs .append (EinsumDense ('abcd,cde->abe' , output_shape = (None , 5 ),
573+ bias_axes = 'e' )(EinsumDense ('abc,cde->abde' ,
574+ output_shape = (None , 3 , 4 ))(inputs [49 ])))
575+
576+ # AdaptiveAvg/MaxPooling 1D/2D/3D variants.
577+ outputs .append (AdaptiveAveragePooling1D (output_size = 3 )(inputs [49 ]))
578+ outputs .append (AdaptiveMaxPooling1D (output_size = 3 )(inputs [49 ]))
579+ outputs .append (AdaptiveAveragePooling1D (output_size = 2 )(inputs [49 ]))
580+ outputs .append (AdaptiveAveragePooling2D (output_size = (3 , 4 ))(inputs [22 ]))
581+ outputs .append (AdaptiveMaxPooling2D (output_size = (3 , 4 ))(inputs [22 ]))
582+ outputs .append (AdaptiveMaxPooling2D (output_size = (13 , 14 ))(inputs [22 ]))
583+ outputs .append (AdaptiveAveragePooling3D (output_size = (2 , 3 , 3 ))(inputs [2 ]))
584+ outputs .append (AdaptiveMaxPooling3D (output_size = (2 , 3 , 3 ))(inputs [2 ]))
585+ outputs .append (AdaptiveAveragePooling3D (output_size = (7 , 2 , 3 ))(inputs [2 ]))
586+
587+ # Discretization on a float input.
588+ outputs .append (Discretization (bin_boundaries = [- 0.5 , 0.0 , 0.5 , 1.0 ])(inputs [49 ]))
589+
590+ # Training-only Random* augmentation layers (passed through at inference).
591+ outputs .append (RandomBrightness (factor = 0.1 )(inputs [22 ]))
592+ outputs .append (RandomFlip (mode = 'horizontal' )(inputs [22 ]))
593+ outputs .append (RandomCrop (height = 26 , width = 28 )(inputs [22 ]))
594+
542595 shared_conv = Conv2D (1 , (1 , 1 ),
543596 padding = 'valid' , name = 'shared_conv' , activation = 'relu' )
544597
@@ -916,89 +969,6 @@ def get_test_model_recurrent() -> Model:
916969 return model
917970
918971
919- def get_test_model_extended () -> Model :
920- """Returns a test model exercising recently-added non-recurrent layers
921- (DepthwiseConv1D, EinsumDense, RMSNormalization, GroupNormalization,
922- AdaptiveAvg/MaxPooling, GroupedQueryAttention, Discretization, and a
923- handful of training-only Random* augmentation passthroughs)."""
924- seq_len = 5
925- n_features = 4
926-
927- inputs = Input (shape = (seq_len , n_features ))
928- inputs_2d = Input (shape = (7 , 8 , 4 ))
929- inputs_3d = Input (shape = (5 , 6 , 7 , 4 ))
930- inputs_kv = Input (shape = (8 , n_features ))
931-
932- dwc1d_same = DepthwiseConv1D (kernel_size = 3 , padding = 'same' )(inputs )
933- dwc1d_valid = DepthwiseConv1D (kernel_size = 2 , padding = 'valid' , strides = 2 )(inputs )
934- dwc1d_no_bias = DepthwiseConv1D (kernel_size = 3 , padding = 'same' , use_bias = False )(inputs )
935-
936- rms = RMSNormalization ()(inputs )
937- rms_eps = RMSNormalization (epsilon = 1e-3 )(inputs )
938- gn = GroupNormalization (groups = 2 )(inputs )
939- gn_no_scale = GroupNormalization (groups = 2 , scale = False )(inputs )
940- gn_no_center = GroupNormalization (groups = 4 , center = False , epsilon = 1e-4 )(inputs )
941- gn_one = GroupNormalization (groups = 1 )(inputs )
942- gn_per_channel = GroupNormalization (groups = n_features )(inputs )
943-
944- ed_dense = EinsumDense ('abc,cd->abd' , output_shape = (None , 8 ), bias_axes = 'd' )(inputs )
945- ed_heads = EinsumDense ('abc,cde->abde' , output_shape = (None , 3 , 4 ), bias_axes = 'de' )(inputs )
946- ed_no_bias = EinsumDense ('abc,cd->abd' , output_shape = (None , 6 ))(inputs )
947- ed_collapse = EinsumDense ('abcd,cde->abe' , output_shape = (None , 5 ), bias_axes = 'e' )(
948- EinsumDense ('abc,cde->abde' , output_shape = (None , 3 , 4 ))(inputs ))
949-
950- ap1d_a = AdaptiveAveragePooling1D (output_size = 3 )(inputs )
951- ap1d_m = AdaptiveMaxPooling1D (output_size = 3 )(inputs )
952- ap1d_alt = AdaptiveAveragePooling1D (output_size = 2 )(inputs )
953- ap2d_a = AdaptiveAveragePooling2D (output_size = (3 , 4 ))(inputs_2d )
954- ap2d_m = AdaptiveMaxPooling2D (output_size = (3 , 4 ))(inputs_2d )
955- ap2d_uneven = AdaptiveMaxPooling2D (output_size = (7 , 2 ))(inputs_2d )
956- ap3d_a = AdaptiveAveragePooling3D (output_size = (2 , 3 , 3 ))(inputs_3d )
957- ap3d_m = AdaptiveMaxPooling3D (output_size = (2 , 3 , 3 ))(inputs_3d )
958- ap3d_uneven = AdaptiveAveragePooling3D (output_size = (5 , 2 , 3 ))(inputs_3d )
959-
960- gqa = GroupQueryAttention (head_dim = 4 , num_query_heads = 6 , num_key_value_heads = 2 )(
961- inputs , inputs , inputs )
962- gqa_gated = GroupQueryAttention (head_dim = 4 , num_query_heads = 6 , num_key_value_heads = 2 ,
963- use_gate = True )(inputs , inputs , inputs )
964- # num_query_heads == num_key_value_heads is the multi-head-attention degenerate case.
965- gqa_eq_heads = GroupQueryAttention (head_dim = 4 , num_query_heads = 4 , num_key_value_heads = 4 )(
966- inputs , inputs , inputs )
967- # distinct sequence lengths for query vs key/value.
968- gqa_kv_diff = GroupQueryAttention (head_dim = 3 , num_query_heads = 4 , num_key_value_heads = 2 )(
969- inputs , inputs_kv , inputs_kv )
970-
971- discretized = Discretization (bin_boundaries = [- 0.5 , 0.0 , 0.5 , 1.0 ])(inputs )
972- rand_brightness = RandomBrightness (factor = 0.1 )(inputs_2d )
973- rand_flip = RandomFlip (mode = 'horizontal' )(inputs_2d )
974- rand_crop = RandomCrop (height = 7 , width = 8 )(inputs_2d )
975-
976- outputs = [
977- dwc1d_same , dwc1d_valid , dwc1d_no_bias ,
978- rms , rms_eps ,
979- gn , gn_no_scale , gn_no_center , gn_one , gn_per_channel ,
980- ed_dense , ed_heads , ed_no_bias , ed_collapse ,
981- ap1d_a , ap1d_m , ap1d_alt ,
982- ap2d_a , ap2d_m , ap2d_uneven ,
983- ap3d_a , ap3d_m , ap3d_uneven ,
984- gqa , gqa_gated , gqa_eq_heads , gqa_kv_diff ,
985- discretized ,
986- rand_brightness , rand_flip , rand_crop ,
987- ]
988-
989- model = Model (inputs = [inputs , inputs_2d , inputs_3d , inputs_kv ],
990- outputs = outputs , name = 'test_model_extended' )
991- model .compile (loss = 'mse' , optimizer = 'adam' )
992-
993- training_data_size = 2
994- data_in = generate_input_data (training_data_size ,
995- [(seq_len , n_features ), (7 , 8 , 4 ), (5 , 6 , 7 , 4 ), (8 , n_features )])
996- initial_data_out = model .predict (data_in )
997- data_out = generate_output_data (training_data_size , initial_data_out )
998- model .fit (data_in , data_out , epochs = 1 )
999- return model
1000-
1001-
1002972def main () -> None :
1003973 """Generate different test models and save them to the given directory."""
1004974 if len (sys .argv ) != 3 :
@@ -1015,7 +985,6 @@ def main() -> None:
1015985 'autoencoder' : get_test_model_autoencoder ,
1016986 'sequential' : get_test_model_sequential ,
1017987 'recurrent' : get_test_model_recurrent ,
1018- 'extended' : get_test_model_extended ,
1019988 }
1020989
1021990 if not model_name in get_model_functions :
0 commit comments