MCPcopy Create free account
hub / github.com/zai-org/CodeGeeX / forward

Method forward

codegeex/megatron/model/transformer.py:164–322  ·  view source on GitHub ↗
(
            self,
            hidden_states,
            attention_mask,
            layer_past=None,
            get_key_value=False,
            prompt_length=None,
            context_length=None,
    )

Source from the content-addressed store, hash-verified

162 skip_bias_add=True)
163
164 def forward(
165 self,
166 hidden_states,
167 attention_mask,
168 layer_past=None,
169 get_key_value=False,
170 prompt_length=None,
171 context_length=None,
172 ):
173 # hidden_states: [sq, b, h]
174
175 # =====================
176 # Query, Key, and Value
177 # =====================
178
179 query_layer, _ = self.query(hidden_states)
180 key_layer, _ = self.key(hidden_states)
181 value_layer, _ = self.value(hidden_states)
182
183 new_query_layer_shape = query_layer.size()[:-1] + \
184 (self.num_attention_heads_per_partition,
185 self.hidden_size_per_attention_head)
186 query_layer = query_layer.view(*new_query_layer_shape)
187
188 new_query_layer_shape = key_layer.size()[:-1] + \
189 (self.num_attention_heads_per_partition,
190 self.hidden_size_per_attention_head)
191 key_layer = key_layer.view(*new_query_layer_shape)
192
193 new_query_layer_shape = value_layer.size()[:-1] + \
194 (self.num_attention_heads_per_partition,
195 self.hidden_size_per_attention_head)
196 value_layer = value_layer.view(*new_query_layer_shape)
197
198 # ==================================
199 # Adjust key and value for inference
200 # ==================================
201
202 if layer_past is not None:
203 past_key, past_value = layer_past
204 key_layer = torch.cat((past_key.type_as(key_layer),
205 key_layer), dim=0)
206 value_layer = torch.cat((past_value.type_as(value_layer),
207 value_layer), dim=0)
208 if get_key_value:
209 present = (key_layer, value_layer)
210
211 # ===================================
212 # Raw attention scores. [b, np, sq, sk]
213 # ===================================
214
215 # [b, np, sq, sk]
216 output_size = (query_layer.size(1),
217 query_layer.size(2),
218 query_layer.size(0),
219 key_layer.size(0))
220
221 # [sq, b, np, hn] -> [sq, b * np, hn]

Callers

nothing calls this directly

Calls 2

sizeMethod · 0.80
forkMethod · 0.80

Tested by

no test coverage detected