Skip to content

Commit 41afc71

Browse files
Dobiasdclaude
andcommitted
Move newly-added non-recurrent layers into test_model_exhaustive
Per review feedback the previous split that introduced a separate test_model_extended duplicated what test_model_exhaustive is for. Drop test_model_extended and append DepthwiseConv1D, EinsumDense, RMSNormalization, GroupNormalization, AdaptiveAvg/MaxPooling, GroupedQueryAttention, Discretization, and the Random* augmentation passthroughs to test_model_exhaustive, reusing existing inputs (notably inputs[49]/[50]/[51] for the attention layers and inputs[22] for image shapes). Co-Authored-By: Claude Opus 4.7 (1M context) <noreply@anthropic.com>
1 parent 2ce984e commit 41afc71

3 files changed

Lines changed: 53 additions & 116 deletions

File tree

keras_export/generate_test_models.py

Lines changed: 53 additions & 84 deletions
Original file line numberDiff line numberDiff line change
@@ -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-
1002972
def 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:

test/CMakeLists.txt

Lines changed: 0 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -28,10 +28,6 @@ add_custom_command ( OUTPUT test_model_recurrent.keras
2828
COMMAND bash -c "${Python3_EXECUTABLE} ${FDEEP_TOP_DIR}/keras_export/generate_test_models.py recurrent test_model_recurrent.keras"
2929
WORKING_DIRECTORY ${CMAKE_BINARY_DIR}/)
3030

31-
add_custom_command ( OUTPUT test_model_extended.keras
32-
COMMAND bash -c "${Python3_EXECUTABLE} ${FDEEP_TOP_DIR}/keras_export/generate_test_models.py extended test_model_extended.keras"
33-
WORKING_DIRECTORY ${CMAKE_BINARY_DIR}/)
34-
3531
add_custom_command ( OUTPUT readme_example_model.keras
3632
COMMAND bash -c "${Python3_EXECUTABLE} ${FDEEP_TOP_DIR}/test/readme_example_generate.py"
3733
WORKING_DIRECTORY ${CMAKE_BINARY_DIR}/)
@@ -66,11 +62,6 @@ add_custom_command ( OUTPUT test_model_recurrent.json
6662
COMMAND bash -c "${Python3_EXECUTABLE} ${FDEEP_TOP_DIR}/keras_export/convert_model.py test_model_recurrent.keras test_model_recurrent.json"
6763
WORKING_DIRECTORY ${CMAKE_BINARY_DIR}/)
6864

69-
add_custom_command ( OUTPUT test_model_extended.json
70-
DEPENDS test_model_extended.keras
71-
COMMAND bash -c "${Python3_EXECUTABLE} ${FDEEP_TOP_DIR}/keras_export/convert_model.py test_model_extended.keras test_model_extended.json"
72-
WORKING_DIRECTORY ${CMAKE_BINARY_DIR}/)
73-
7465
add_custom_command ( OUTPUT readme_example_model.json
7566
DEPENDS readme_example_model.keras
7667
COMMAND bash -c "${Python3_EXECUTABLE} ${FDEEP_TOP_DIR}/keras_export/convert_model.py readme_example_model.keras readme_example_model.json"
@@ -93,7 +84,6 @@ _add_test(test_model_variable_test test_model_variable.json)
9384
_add_test(test_model_autoencoder_test test_model_autoencoder.json)
9485
_add_test(test_model_sequential_test test_model_sequential.json)
9586
_add_test(test_model_recurrent_test test_model_recurrent.json)
96-
_add_test(test_model_extended_test test_model_extended.json)
9787
_add_test(readme_example_main readme_example_model.json)
9888

9989
add_custom_target(unittest
@@ -103,7 +93,6 @@ add_custom_target(unittest
10393
COMMAND test_model_autoencoder_test
10494
COMMAND test_model_sequential_test
10595
COMMAND test_model_recurrent_test
106-
COMMAND test_model_extended_test
10796
COMMAND readme_example_main
10897

10998
COMMENT "Running unittests\n\n"

test/test_model_extended_test.cpp

Lines changed: 0 additions & 21 deletions
This file was deleted.

0 commit comments

Comments
 (0)