MCPcopy Create free account
hub / github.com/google-deepmind/alphageometry / TransformerLayerGenerate

Class TransformerLayerGenerate

transformer_layer.py:55–527  ·  view source on GitHub ↗

Full transformer layer, with attention.

Source from the content-addressed store, hash-verified

53
54@gin.configurable
55class 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[

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected