Full transformer layer, with attention.
| 53 | |
| 54 | @gin.configurable |
| 55 | class TransformerLayerGenerate(transformer_layer.TransformerLayer): |
| 56 | """Full transformer layer, with attention.""" |
| 57 | |
| 58 | def _next_decoder_state( |
| 59 | self, decoder_state: DecoderState, keys: Array, values: Array |
| 60 | ) -> Tuple[DecoderState, Array, Array]: |
| 61 | """Compute the next decoder state, and return keys,values to attend to. |
| 62 | |
| 63 | The keys,values returned from this function are drawn from the prior |
| 64 | decoding state, and comprise a full window of local context. |
| 65 | |
| 66 | Args: |
| 67 | decoder_state: The current decoder state, initially created using |
| 68 | init_decoder_state(). |
| 69 | keys: The key for the current token, of shape (batch_size, 1, dim) |
| 70 | values: The value for the current token of shape (batch_size, 1, dim) |
| 71 | |
| 72 | Returns: |
| 73 | (next_decoder_state, |
| 74 | window of keys of shape (batch_size, window_length, dim), |
| 75 | window of values of shape (batch_size, window_length, dim)) |
| 76 | """ |
| 77 | |
| 78 | assert keys.shape[1] == 1 # single-token autoregressive decoding. |
| 79 | |
| 80 | # Unpack decoder_state |
| 81 | stored_keys = decoder_state["keys"] |
| 82 | stored_values = decoder_state["values"] |
| 83 | curr_index = decoder_state["current_index"] |
| 84 | |
| 85 | # Slice to get window_length-sized chunk of previous keys,values. |
| 86 | out_decoder_state = {} |
| 87 | curr_win_index = curr_index - self.window_length |
| 88 | |
| 89 | # out_keys = jax.lax.dynamic_slice_in_dim( |
| 90 | # stored_keys, curr_win_index, self.window_length, axis=1) |
| 91 | out_keys = slice_in_dim_1(self.window_length)(stored_keys, curr_win_index) |
| 92 | |
| 93 | # out_values = jax.lax.dynamic_slice_in_dim( |
| 94 | # stored_values, curr_win_index, self.window_length, axis=1) |
| 95 | out_values = slice_in_dim_1(self.window_length)( |
| 96 | stored_values, curr_win_index |
| 97 | ) |
| 98 | |
| 99 | # Write current keys,values to stored keys, values. |
| 100 | # stored_keys = jax.lax.dynamic_update_slice_in_dim( |
| 101 | # stored_keys, keys, curr_index, axis=1) |
| 102 | stored_keys = update_slice_in_dim_1(stored_keys, keys, curr_index) |
| 103 | # stored_values = jax.lax.dynamic_update_slice_in_dim( |
| 104 | # stored_values, values, curr_index, axis=1) |
| 105 | stored_values = update_slice_in_dim_1(stored_values, values, curr_index) |
| 106 | curr_index = curr_index + 1 |
| 107 | |
| 108 | # Pack a new decoder_state object. |
| 109 | out_decoder_state["keys"] = stored_keys |
| 110 | out_decoder_state["values"] = stored_values |
| 111 | out_decoder_state["current_index"] = curr_index |
| 112 | out_decoder_state["relative_position_bias"] = decoder_state[ |
nothing calls this directly
no outgoing calls
no test coverage detected