|
42 | 42 | from keras.layers import RNN, LSTMCell, GRUCell, SimpleRNNCell, StackedRNNCells |
43 | 43 | from keras.layers import Discretization, IntegerLookup, Masking |
44 | 44 | from keras.layers import RandomBrightness, RandomFlip, RandomCrop |
| 45 | +from keras.layers import RandomContrast, RandomRotation, RandomTranslation, RandomZoom |
| 46 | +from keras.layers import RandomHue, RandomSaturation, RandomSharpness |
45 | 47 | from keras.models import Model, load_model, Sequential |
46 | 48 |
|
47 | 49 | __author__ = "Tobias Hermann" |
@@ -546,6 +548,10 @@ def get_test_model_exhaustive() -> Model: |
546 | 548 | num_key_value_heads=2, use_gate=True)(inputs[49], inputs[49], inputs[49])) |
547 | 549 | outputs.append(GroupQueryAttention(head_dim=4, num_query_heads=4, |
548 | 550 | num_key_value_heads=4)(inputs[49], inputs[49], inputs[49])) |
| 551 | + # num_query_heads == num_key_value_heads with distinct K/V seq lengths |
| 552 | + # exercises the MHA-equivalence path with cross-attention shapes. |
| 553 | + outputs.append(GroupQueryAttention(head_dim=4, num_query_heads=4, |
| 554 | + num_key_value_heads=4)(inputs[49], inputs[50], inputs[51])) |
549 | 555 | outputs.append(GroupQueryAttention(head_dim=3, num_query_heads=4, |
550 | 556 | num_key_value_heads=2)(inputs[49], inputs[50], inputs[51])) |
551 | 557 |
|
@@ -587,10 +593,21 @@ def get_test_model_exhaustive() -> Model: |
587 | 593 | # Discretization on a float input. |
588 | 594 | outputs.append(Discretization(bin_boundaries=[-0.5, 0.0, 0.5, 1.0])(inputs[49])) |
589 | 595 |
|
| 596 | + # IntegerLookup: scaled to int via Discretization so the input is float. |
| 597 | + outputs.append(IntegerLookup(vocabulary=[0, 1, 2, 3])( |
| 598 | + Discretization(bin_boundaries=[-1.0, 0.0, 1.0])(inputs[49]))) |
| 599 | + |
590 | 600 | # Training-only Random* augmentation layers (passed through at inference). |
591 | 601 | outputs.append(RandomBrightness(factor=0.1)(inputs[22])) |
592 | 602 | outputs.append(RandomFlip(mode='horizontal')(inputs[22])) |
593 | 603 | outputs.append(RandomCrop(height=26, width=28)(inputs[22])) |
| 604 | + outputs.append(RandomContrast(factor=0.1)(inputs[22])) |
| 605 | + outputs.append(RandomRotation(factor=0.1)(inputs[22])) |
| 606 | + outputs.append(RandomTranslation(0.1, 0.1)(inputs[22])) |
| 607 | + outputs.append(RandomZoom(0.1)(inputs[22])) |
| 608 | + outputs.append(RandomHue(factor=0.1, value_range=(0.0, 1.0))(inputs[22])) |
| 609 | + outputs.append(RandomSaturation(factor=0.1, value_range=(0.0, 1.0))(inputs[22])) |
| 610 | + outputs.append(RandomSharpness(factor=0.1, value_range=(0.0, 1.0))(inputs[22])) |
594 | 611 |
|
595 | 612 | shared_conv = Conv2D(1, (1, 1), |
596 | 613 | padding='valid', name='shared_conv', activation='relu') |
@@ -901,6 +918,9 @@ def get_test_model_recurrent() -> Model: |
901 | 918 | rnn_seq = SimpleRNN(6, return_sequences=True)(inputs) |
902 | 919 | rnn_last = SimpleRNN(5, activation='tanh')(rnn_seq) |
903 | 920 | rnn_relu_no_bias = SimpleRNN(4, activation='relu', use_bias=False)(inputs) |
| 921 | + rnn_state_out, rnn_state_h = SimpleRNN(3, return_state=True)(inputs) |
| 922 | + |
| 923 | + gru_state_out, gru_state_h = GRU(4, return_state=True)(inputs) |
904 | 924 |
|
905 | 925 | bidi_lstm_seq = Bidirectional(LSTM(4, return_sequences=True))(inputs) |
906 | 926 | bidi_lstm_last_concat = Bidirectional(LSTM(3))(bidi_lstm_seq) |
@@ -941,8 +961,8 @@ def get_test_model_recurrent() -> Model: |
941 | 961 | outputs = [ |
942 | 962 | lstm_last, lstm_no_bias, lstm_relu, lstm_no_unit_forget, |
943 | 963 | lstm_state_out, lstm_state_h, lstm_state_c, |
944 | | - gru_last, gru_no_bias, |
945 | | - rnn_last, rnn_relu_no_bias, |
| 964 | + gru_last, gru_no_bias, gru_state_out, gru_state_h, |
| 965 | + rnn_last, rnn_relu_no_bias, rnn_state_out, rnn_state_h, |
946 | 966 | bidi_lstm_last_concat, |
947 | 967 | bidi_lstm_sum, |
948 | 968 | bidi_gru_mul, |
|
0 commit comments